jax.experimental.pallas.triton 模块

目录

jax.experimental.pallas.triton 模块#

Triton 特定的 Pallas API。

#

CompilerParams([num_warps, num_stages])

Triton 的编译器参数。

函数#

atomic_and(x_ref_or_view, idx, val, *[, mask])

原子性计算 x_ref_or_view[idx] &= val

atomic_add(x_ref_or_view, idx, val, *[, mask])

原子性计算 x_ref_or_view[idx] += val

atomic_cas(ref, cmp, val)

对引用中的值执行原子比较并交换(Compare-and-Swap)操作,

atomic_max(x_ref_or_view, idx, val, *[, mask])

原子性计算 x_ref_or_view[idx] = max(x_ref_or_view[idx], val)

atomic_min(x_ref_or_view, idx, val, *[, mask])

原子性计算 x_ref_or_view[idx] = min(x_ref_or_view[idx], val)

atomic_or(x_ref_or_view, idx, val, *[, mask])

原子性计算 x_ref_or_view[idx] |= val

atomic_xchg(x_ref_or_view, idx, val, *[, mask])

将给定值与指定索引处的值进行原子交换。

atomic_xor(x_ref_or_view, idx, val, *[, mask])

原子性计算 x_ref_or_view[idx] ^= val

approx_tanh(x)

逐元素近似双曲正切函数:\(\mathrm{tanh}(x)\)

debug_barrier()

同步网格中的所有内核执行。

elementwise_inline_asm(asm, *, args, ...)

应用逐元素操作的内联汇编。

load(ref, *[, mask, other, cache_modifier, ...])

从给定的引用加载数组。

max_contiguous(x, values)

编译器提示,断言 x 的前 values 个值是连续的。

store(ref, val, *[, mask, eviction_policy])

将值存储到给定的引用中。