【问题标题】:Conditioning on elements of matrix in JIT-ted functionJIT-ted函数中矩阵元素的条件化
【发布时间】:2021-10-16 21:52:39
【问题描述】:

我有一个看起来像这样的函数

        @jax.jit
        def f(R):
            tr = jnp.trace(R)

            r00 = R[0, 0]
            r01 = R[0, 1]
            r02 = R[0, 2]
            r10 = R[1, 0]
            r11 = R[1, 1]
            r12 = R[1, 2]
            r20 = R[2, 0]
            r21 = R[2, 1]
            r22 = R[2, 2]

            condw = tr > 0
            condx = (r00 > r11) and (r00 > r22)
            condy = (r11 > r22)
            # ... do some more things based on the conditions

其中R 是一个 3x3 DeviceArray。当我尝试 JIT-ing 这个函数时,如上所示,我收到以下错误:

File "/path/to/my/file", line 90, in f
    condx = (r00 > r11) and (r00 > r22)
  File "/Users/me/miniconda3/envs/myenv/lib/python3.9/site-packages/jax/core.py", line 544, in __bool__
    def __bool__(self): return self.aval._bool(self)
  File "/Users/me/miniconda3/envs/myenv/lib/python3.9/site-packages/jax/core.py", line 989, in error
    raise ConcretizationTypeError(arg, fname_context)
jax._src.traceback_util.UnfilteredStackTrace: jax._src.errors.ConcretizationTypeError: Abstract tracer value encountered where concrete value is expected: Traced<ShapedArray(bool[])>with<DynamicJaxprTrace(level=0/1)>
The problem arose with the `bool` function. 
While tracing the function f at /path/to/my/file:75 for jit, this concrete value was not available in Python because it depends on the value of the argument 'R'.

我不太确定计算这个布尔值会阻止这个函数被 JIT-ted 有什么问题。

        condx = (r00 > r11) and (r00 > r22)

任何提示将不胜感激 - 谢谢!

【问题讨论】:

    标签: python jax


    【解决方案1】:

    #3761开始,使用位运算符而不是逻辑运算符。

    这行得通。

    condx = (r00 > r11) & (r00 > r22)
    

    【讨论】:

      猜你喜欢
      • 1970-01-01
      • 1970-01-01
      • 2012-08-12
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 1970-01-01
      • 2021-08-04
      相关资源
      最近更新 更多