- 机器学习
- 深度学习
【免费下载链接】jax
Composable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more
导读:本文是 docs/pallas/tpu 系列文档的深度整理,主题为如何使用 Pallas(JAX 的自定义内核语言)在 Google TPU 上编写高性能内核。文章先剖析 TPU 与 GPU 截然不同的硬件特性(HBM/VMEM/SMEM 内存层级、向量寄存器、顺序执行语义),再逐步讲解
BlockSpec/grid迭代规则、dimension_semantics多核并行、SMEM 标量预取,以及用grid+BlockSpec自动生成流水线来重叠内存拷贝与计算。读完你将掌握 Pallas TPU 内核的正确编写姿势、性能调优规则和常见陷阱,并能直接复现文中的矩阵加法、归约求和等完整内核示例。
Pallas TPU 后端:实验状态与正确性承诺
Pallas 是 JAX 的扩展,用于为 GPU 与 TPU 编写自定义内核(见 docs/pallas/index.rst 的定位说明)。其中 TPU 后端由 Mosaic 编译管线支撑,相关 API 集中在 jax/experimental/pallas/tpu.py(VMEM、SMEM、PrefetchScalarGridSpec、emit_pipeline等均由此导出)。
使用 TPU 后端前需要明确两点:
- 实验性:该功能仍处于实验阶段,目前只接受 JAX NumPy 的一个子集,且官方仍在改进错误信息。因此编写内核时遇到 "not implemented" 错误并不罕见。
- 正确性承诺:虽然功能实验性,但官方对正确性非常严肃——只要内核被编译器接受,就必然返回预期结果。
当发现意外输出时,务必用interpret=True传入pallas_call再跑一遍同一内核进行对照。从源码看,interpret模式会把pallas_call实现为对grid的jax.jit扫描(见 jax/_src/pallas/pallas_call.py 中 "runs thepallas_callas ajax.jitof a scan over the grid" 的说明),它不需要 TPU 或 GPU,是唯一能在 CPU 上运行 Pallas 内核的方式,非常适合调试。若两种模式结果不一致,则属于编译器 bug,应提交 bug report。
理解 TPU:与 GPU 截然不同的硬件
内存空间、寄存器与计算单元
TPU 是 Google 开发的专用机器学习加速器。与 GPU 相比,其核心差异在于:TPU 是带有超宽向量寄存器的顺序执行机器(类似 CPU),同时允许软件把部分操作调度到后台异步执行。
TPU 的 TensorCore 由三类构件组成(示意图见 docs/pallas/tpu/pipelining.md):
- 内存空间:
- HBM(高带宽内存):常被理解为"设备内存",数组的主要存放地;
- VMEM(向量内存):用于存放向量/数组值的缓存;
- SMEM(标量内存):用于存放标量值的缓存。
- 寄存器:向量寄存器(VREGs)存数组值,标量寄存器(SREGs)存标量值;数据从对应缓存加载进寄存器。
- 计算单元:标量单元(scalar unit)、向量单元(VPU)、矩阵单元(MXU)。计算单元只对寄存器中的值运算,结果也写回寄存器。
异步后台操作
TPU 允许软件调度以下操作在后台执行,使其与主指令流异步:
- HBM 内存访问:不能直接发起,必须由 DMA 子单元预取到更低层级内存;
- 矩阵乘法:由 MXU 单元支持;
- 矩阵转置与置换:由 XLU 单元支持。
一次向量化计算的完整链路
对 HBM 中的x、y做向量加法,硬件层面需要五步:
- 把
x、y从 HBM 拷贝到 VMEM; - 从 VMEM 加载到 VREGs;
- 用 VPU 或 MXU 执行计算,结果写入 VREGs;
- 把输出 VREGs 存回 VMEM;
- 把 VMEM 中的输出拷回 HBM。
BlockSpec 与 grid 迭代:Pallas TPU 内核的核心约束
BlockSpec(block_shape+index_map,定义于 jax/_src/pallas/core.py)与grid的行为在 TPU 上大体符合 Pallas 的一般语义:每次调用内核体,都拿到输入的切片并负责初始化输出切片。但 TPU 后端有几条特殊规则:
窗口形状限制
不是所有窗口形状都受支持。如果输入的最后两维分别大于 8 和 128,那么这两个维度的窗口形状必须是 8 和 128 的倍数;如果输入维度更小,窗口应覆盖整个维度。
内存空间:HBM 与 VMEM/SMEM 的分工
pallas_call的输入通常位于 HBM,但传入内核体的引用(Ref)指向的是更低层内存(VMEM 或 SMEM)。这让内核体能以极高速读写它们,而所有与 HBM 的通信(延迟极高)由编译器负责,并与计算重叠。这正是 Pallas TPU 内核高性能的来源。
顺序 grid 语义带来的三个推论
与 GPU 不同,TPU 高度顺序化,grid 通常按字典序顺序执行而非并行(Multicore 配置除外,见后文)。这带来三个重要能力:
- HBM 传输复用:当两个(字典序)相邻的 grid 索引使用同一输入切片时,第二次迭代的 HBM 传输会被跳过——数据已就绪;
- 无竞争输出写入:多次内核体调用可以写同一输出切片而不会产生竞争条件。但要求:所有写同一切片的调用必须连续;
- 输出切片"前缀-后缀"结构:输出的"连续"限制通常意味着:grid 维度的某个前缀总在变化输出切片,而输出窗口对剩余后缀保持恒定。
以矩阵乘法内核为例:一般用 3 维 grid——前两维分别对应左操作数第一轴、右操作数第二轴的切片,第三个(最后)轴负责 tile reduction 维度。reduction 轴必须是最后一维,因为输出窗口在该轴上不变化,输出引用因此可以作为部分和累加器反复使用。
VMEM 容量与窗口大小
VMEM 对这么底层的存储层级来说相当大(16MB+),所以可以用很大的窗口。经验上窗口越大,硬件利用率通常越好。但如果窗口(加上溢出的向量寄存器所需空间)超过 VMEM 容量,就会看到底层编译器报出的内存不足(OOM)错误。
维度顺序有意义:把最后两维做大
在普通jax.jit程序中,中间数组的维度顺序通常不影响性能——编译器可自由重排。但 Pallas 暴露的是底层能力,维度顺序对生成代码质量影响巨大。
原因在于:TPU 的绝大部分计算在 2D 向量寄存器上进行,而Pallas TPU 只会把中间数组的最后两维映射到向量寄存器维度(分别对应 sublanes 和 lanes)。形状为(n, 1, 1)的数组至少需要n个向量寄存器表示;n过大时可能导致寄存器溢出和因内存占用过大而触发 VMEM OOM。虽然底层编译器很擅长重排指令以降低寄存器压力,但稳妥的经验法则是:
保持最后两维(尤其是最后一维)较大,前导维度较小。
Multicore TPU 配置:dimension_semantics 与 Megacore
单芯片双核抽象
在较新的 TPU 世代(如 v4、v5p),芯片上的两个 TensorCore 常被抽象为单一设备(即Megacore模式)。两个 TensorCore 各有独立的 VMEM、VREGs、SMEM、SREGs 和计算单元,但共享 HBM。概念上,Megacore 设备像一台只有两个线程的极简 GPU。
要利用多核,Pallas 必须打破顺序 grid 执行保证,把某个 grid 轴并行化到各核上。这是**显式选择(opt-in)**的,通过pallas_call的compiler_params传入dimension_semantics实现:
pallas_call( ..., compiler_params=dict( mosaic=dict( dimension_semantics=["parallel", "parallel", "arbitrary"] ) ), )dimension_semantics 的语义
该参数是一个列表,条目数与 grid 轴数相同。只有标记为"parallel"的维度可以在核间分区。经验法则:输出窗口不变的维度才是 parallel 的(否则无法无竞争地并行);因此dimension_semantics永远是若干个parallel轴后跟若干个arbitrary轴。"arbitrary"表示该维度不可做任何假设,因此不能并行化。
从源码看,该参数在 Mosaic 管线中经过严格校验:未提供时默认全为"arbitrary",长度必须与grid一致(见 jax/_src/pallas/mosaic/lowering.py 与 jax/_src/pallas/mosaic/pipeline.py),最终以dimension_semantics属性附着到编译产物上。在pallas_call_registration.py中,它从mosaic_params中提取并传入 lowering(见 jax/_src/pallas/mosaic/pallas_call_registration.py)。
在流水线示例中启用 Megacore 只需一行注解:
def add_matrices_pipelined_megacore(x: jax.Array, y: jax.Array) -> jax.Array: block_spec = pl.BlockSpec((256, 512), lambda i: (i, 0)) return pl.pallas_call( add_matrices_kernel, out_shape=x, in_specs=[block_spec, block_spec], out_specs=block_spec, grid=(2,), compiler_params=dict(mosaic=dict(dimension_semantics=("parallel",))) )(x, y)指定dimension_semantics后,Pallas 会自动把 grid 拆分并在两个 TensorCore 上同时执行。
注意:Megacore 目前仅对 TPU v4 和 v5p 生效。在其他平台上提供该注解是 no-op,但不指定它会导致即使有多个核也只用一个 TensorCore。
并行化的收益与风险
在 2 核 TPU 上分区内核常带来约 2 倍加速,但实际收益可能显著小于 2 倍——尤其是当不同内核体实例的计算成本差异很大时:若所有昂贵步骤恰好被映射到一个核、廉价步骤全在另一个核,第二个核会空转到第一个核完成。此外,Pallas TPU 一般偏好分区大小为核数倍数的轴,并优先分区前导 grid 轴。
把操作数放进 SMEM:PrefetchScalarGridSpec
TPU 上大部分计算发生在向量单元,但控制流等场景常需要标量运算。为此 TPU 配有独立的标量单元和标量内存(SMEM)。经验法则:任何用于控制流决策的数据都应放在 SMEM。
SMEM 是低延迟内存,支持随机访问,但单条指令只能读写 32 位值——相比 VMEM 事务 4KBi 的粒度小得多,却因为没有对齐要求而灵活得多。
当内核不以规则模式访问输入 tile 时(例如块稀疏内核),标量内存非常有用。在 Pallas 中,把grid参数换成grid_spec为PrefetchScalarGridSpec、并设置非零的num_scalar_prefetch即可实现。规则如下:
- 若
num_scalar_prefetch为n,则pallas_call的前n个参数被放入 SMEM; - 这些参数不应指定
BlockSpec; - 后续所有参数的
BlockSpec的index_map会额外收到这些 SMEM 引用(即 index_map 签名中多了前导标量参数)。
从源码看,PrefetchScalarGridSpec定义于 jax/_src/pallas/mosaic/core.py:它继承自GridSpec,构造参数为num_scalar_prefetch、grid、in_specs、out_specs、scratch_shapes;在get_grid_mapping中,前num_scalar_prefetch个输入被切分出来并映射为TPUMemorySpace.SMEM的引用(jax/_src/pallas/mosaic/core.py)。
配套的测试用例(如 tests/pallas/tpu/pallas_call_test.py)展示了标准用法:标量索引s被预取进 SMEM,x的index_map接收(i, s_ref)并在内核内用pl.load(s_ref, (i,))读取标量来决定切片位置。该类测试均可在interpret=True下于 CPU 验证。
支持的数据类型与计算放置
目前 Pallas TPU仅支持以下数据类型:
| 类别 | 支持情况 |
|---|---|
jnp.float32 | 支持 |
jnp.bfloat16 | 支持 |
jnp.int* | 支持所有精度,除jnp.int4 |
jnp.uint* | 支持所有精度 |
计算放置规则:所有标量(0D)数组存放在标量寄存器中,相关运算在标量核执行;其余所有操作(即使是单元素的 1D+ 数组)都在向量核执行。
支持的操作全景与性能特征
矩阵乘法
- 矩阵乘法结果始终是 float32。若输入不是 float32,推荐用
lax.dot并设preferred_element_type=jnp.float32; - 使用
lax.dot_general时,操作数最后两维的转置可以融合进乘法,提升整体性能。
精度控制:Pallas TPU 的 lowering 感知jax.default_matmul_precision。追求最高性能(和最低精度)用bfloat16;关心数值精度则设为float32。
警告:即使给矩阵乘法传入 32 位操作数,除非显式请求
float32精度,它们仍会被舍入到bfloat16。
转置
- 值至少有 4 维时,除最后两维外的任意轴转置是免费的;
- 否则只实现了最后两维的转置;
- 注意:最后两维的某些转置可以融合进矩阵乘法。
内存访问
引用(Ref)的任意切片均可读可写,但要受实现约束:
- 32 位宽的输入目前没有限制;
- 更窄的类型只支持部分切片模式;
- 最后两维上对齐到 8 和 128 的倍数、且长度为 8 和 128 的倍数的读写总是受支持。
由于向量内存的读写通常发生在(8, 128)的 tile 上,读写至少二维的引用时,最佳性能条件是:基址偏移可被 tiling 整除,读取区域大小是 tile 大小的倍数。
元素操作
硬件一般只支持用32 位类型做逐元素计算。加载低精度操作数时,通常应先 upcast 到 32 位类型再做元素操作。不同元素操作的成本差异显著,官方将其分为三档:
| 操作 | 成本 |
|---|---|
jnp.add、+ | 🟢 便宜 |
jnp.sub、- | 🟢 便宜 |
jnp.mul、* | 🟢 便宜 |
/、//、% | 🌕 中等 |
jnp.max、jnp.min | 🟢 便宜 |
jnp.where(select) | 🟢 便宜 |
jnp.abs | 🟢 便宜 |
\|、^、&、~ | 🟢 便宜 |
<<、>> | 🟢 便宜 |
比较(==等) | 🟢 便宜 |
类型转换(.astype) | 🟢 便宜 |
jnp.exp | 🌕 中等 |
jnp.tanh | 🌕 中等 |
jnp.pow | 🌕 中等 |
jnp.sin | 🔴 昂贵 |
jnp.cos | 🔴 昂贵 |
许多 JAX 函数由其他原语组合实现,故该表并非穷尽。例如jax.nn.relu由比较和jnp.where实现,因此也能在 Pallas 内核中工作。
数组构造器
所有常量数组构造器均受支持(jnp.ones、jnp.zeros、jnp.full)。值得注意:jax.random模块目前与 Pallas 不兼容。
归约
sum、max、min归约受支持,但每次只能归约单个数组轴,且性能差异明显:
- 最后一维归约一般最慢;
- 倒数第二维归约较快,但仍慢于前导维。
广播
广播的性能特征与归约非常相似:
- 除最后两维外的广播总是受支持且免费;
- 沿倒数第二维广播较慢;
- 沿最后一维广播最慢。
Reshape
- 除最后两维外的 reshape 均受支持且免费;
- 能修改最后两维的 reshape 只有两种受支持情况:(1)某些前导维被展平到倒数第二维上;(2)添加一个刚被归约移除的维度。
控制流
TPU 后端目前对控制流支持有限,支持cond、fori_loop和for_loop。但循环原语在编译期会被完全展开,所以请把循环次数(trip count)控制在合理的小范围内。过度使用控制流会导致底层代码生成显著退化,推荐尽可能把计算密集的操作挤进单个基本块。
用流水线重叠内存 I/O 与计算
VMEM/SMEM 的两大约束
使用pallas_call直接搬运数组有两大约束:
- 容量:VMEM 和 SMEM 很小!v4 TPU 的 VMEM 只有 16MiB,SMEM 只有几十到几百 KiB。作为参照,一个
f32[2048, 2048]数组恰好 16MiB——上面的朴素内核无法扩展到超出中等规模的数组; - 带宽:HBM↔VMEM 的拷贝比绝大多数计算指令慢得多。
add_matrices大概率花在 HBM/VMEM 间拷贝上的时间多于加法本身。
流水线的思想:把拷贝藏进计算
流水线的目标是在并行地做 HBM↔VMEM 拷贝的同时利用计算单元。朴素程序的问题在于:先把所有x、y拷贝完才开始计算,造成拷贝与计算的串行依赖。
如果把计算切分成多个子计算(例如把矩阵加法拆成若干"块"的加法),就能用一块子计算的拷贝去重叠另一块子计算的计算。以把(512, 512)的x、y沿前导轴切成x1, x2、y1, y2(各(256, 512))为例,流水线执行序列为:
- 拷贝
x1、y1进 VMEM; - 开始(异步)拷贝
x2、y2进 VMEM; - 从 VMEM 加载
x1, y1到 VREGs; - 计算
z1 = x1 + y1; - 把
z1存入 VMEM; - 开始把
z1从 VMEM 拷回 HBM; - 等待
x2, y2拷贝完成; - 加载
x2, y2到 VREGs; - 计算
z2 = x2 + y2; - 存
z2进 VMEM;等待z1拷完;开始拷z2回 HBM;最后等待z2拷完。
任何时刻只要在做计算,就有异步拷贝在进行——拷贝的部分时间没有被浪费。
判定流水线效率的两个关键数字:需要执行的浮点运算量(FLOPs)和为此需要拷贝的字节数。二者之比(FLOPs/内存字节)称为算术强度(arithmetic intensity),它决定了流水线是计算受限还是内存受限。
用 grid 与 BlockSpec 表达流水线
Pallas 用grid和BlockSpec自动生成上述流水线,无需手写异步序列。在流水线中,BlockSpec.block_shape为(256, 512),第一次迭代取x1、第二次取x2:
def x_index_map(i): return (i, 0) block_spec = pl.BlockSpec((256, 512), x_index_map)完整内核如下——pallas_call负责把x、y拷入 VMEM、分配输出 VMEM 缓冲区(z_vmem_ref),并在内核结束后把输出拷回 HBM:
def add_matrices_kernel(x_vmem_ref, y_vmem_ref, z_vmem_ref): # 从 VMEM 加载到 VREGs x_vregs = x_vmem_ref[:, :] y_vregs = y_vmem_ref[:, :] # 执行向量加法 z_vregs = x_vregs + y_vregs # 把 VREGs 中的结果存回 VMEM z_vmem_ref[:, :] = z_vregs def add_matrices_pipelined(x: jax.Array, y: jax.Array) -> jax.Array: block_spec = pl.BlockSpec((256, 512), lambda i: (i, 0)) return pl.pallas_call( add_matrices_kernel, out_shape=x, in_specs=[block_spec, block_spec], out_specs=block_spec, grid=(2,) )(x, y)只加了很少的代码,BlockSpec和grid就完成了大量工作:BlockSpec提供了足够信息去预取输入块——例如迭代i时,把i + 1传入index_map得到下一迭代所需的块,然后发起异步拷贝;对输出则等待上一迭代输出拷贝完成,再开始当前迭代输出的拷贝。pallas_call的完整签名(含grid_spec、debug、interpret、compiler_params等参数)见 jax/_src/pallas/pallas_call.py。
参数化流水线:块大小是最重要的调优旋钮
块大小是优化 Pallas 内核性能时最重要的调优参数:更小的块会给流水线循环增加更多迭代,每次迭代工作量更少。还可以同时沿第二维切分输入输出:
def add_matrices_pipelined_2d( x: jax.Array, y: jax.Array, *, bm: int = 256, bn: int = 256 ) -> jax.Array: m, n = x.shape block_spec = pl.BlockSpec((bm, bn), lambda i, j: (i, j)) return pl.pallas_call( add_matrices_kernel, out_shape=x, in_specs=[block_spec, block_spec], out_specs=block_spec, grid=(m // bm, n // bn), )(x, y) np.testing.assert_array_equal( add_matrices_pipelined_2d(x, y, bm=256, bn=256), x + y ) np.testing.assert_array_equal( add_matrices_pipelined_2d(x, y, bm=128, bn=128), x + y ) np.testing.assert_array_equal( add_matrices_pipelined_2d(x, y, bm=512, bn=512), x + y )二维grid在底层被降级为嵌套循环,即多层流水线。
处理归约:初始化累加器是关键
把(8, 512, 512)的数组沿轴 0 归约为(512, 512),可用大小为(8,)的grid,每次迭代把x[i]累加到输出 VMEM 缓冲区。朴素实现是错误的:
# Warning: this implementation is incorrect! def naive_sum_kernel(x_ref, o_ref): o_ref[...] += x_ref[...] def naive_sum(x: jax.Array) -> jax.Array: grid, *out_shape = x.shape return pl.pallas_call( naive_sum_kernel, grid=grid, # block_shape 中的 None 表示取大小为 1 并在内核中挤掉该维 in_specs=[pl.BlockSpec((None, *out_shape), lambda i: (i, 0, 0))], out_specs=pl.BlockSpec(out_shape, lambda i: (0, 0)), out_shape=jax.ShapeDtypeStruct(out_shape, x.dtype), )(x)这里in_specs把(512, 512)整块载入 VMEM(该维度不流水线化),index_map每次选x的第i维;block_shape中的None表示选择一个单例维度并在内核里挤掉,因此内核里的x_ref是(512, 512)。out_specs的index_map为lambda i: (0, 0),说明o_ref在流水线中保持不变,每次迭代都可读写它。
问题在于:o_ref初始内容是垃圾,累加会基于垃圾进行,导致结果错误。因此:
在内核中做归约时,必须初始化保存归约值的
Ref。
用pl.when(jax.lax.cond的便捷封装)配合pl.program_id(查询当前 grid 轴的迭代序号)在迭代 0 初始化:
def sum_kernel(x_ref, o_ref): @pl.when(pl.program_id(axis=0) == 0) def _(): o_ref[...] = jnp.zeros_like(o_ref) o_ref[...] += x_ref[...] def sum(x: jax.Array) -> jax.Array: grid, *out_shape = x.shape return pl.pallas_call( sum_kernel, grid=grid, in_specs=[pl.BlockSpec((None, *out_shape), lambda i: (i, 0, 0))], out_specs=pl.BlockSpec(out_shape, lambda i: (0, 0)), out_shape=jax.ShapeDtypeStruct(out_shape, x.dtype) )(x)归约必须放在 grid 的最右(minormost)维度
Pallas 用BlockSpec、grid和内核函数生成的流水线不会从 HBM 读回输出——输出一旦写回 HBM 就无法再访问。因此,不能跨越"会被重新访问"的 grid 维度做归约,所有归约都必须发生在 grid 的最右维度。(上例 grid 只有一维,天然满足。)
排查与验证要点
- interpret 模式优先:所有内核先用
interpret=True在 CPU 上验证逻辑正确性(源码文档明确这是"唯一的 CPU 运行方式"),再上真机。测试目录 tests/pallas/tpu/pallas_call_test.py 中的用例(如标量预取、vmap 组合)均以interpret参数双路径覆盖,可作为编写参考; - 对齐规则:涉及向量内存访问时,优先保证最后两维偏移可被
(8, 128)tile 整除、长度是其倍数; - VMEM 预算:窗口过大导致的 VMEM OOM 会以底层编译器错误出现,注意评估窗口 + 寄存器溢出所需空间;
- 控制流预算:循环会被完全展开,保持 trip count 小,把计算密集操作放进同一基本块。
总结
Pallas 让开发者在不必完全理解 TPU 硬件的情况下也能开始写内核,但理解硬件显然更利于写出高性能内核。本文覆盖了 Pallas TPU 的核心知识体系:硬件内存模型(HBM/VMEM/SMEM)与顺序执行语义、BlockSpec/grid的切片与连续输出规则、维度顺序对寄存器压力的影响、dimension_semantics多核(Megacore)并行、PrefetchScalarGridSpec的 SMEM 标量预取、支持的数据类型与各类操作的成本特征,以及用grid+BlockSpec自动生成流水线来重叠内存 I/O 与计算、并正确处理归约累加器初始化。
文中内核均可在 docs/pallas/tpu/pipelining.md 中查看完整可运行版本。延伸阅读:grid与BlockSpec的通用概念见 docs/pallas/grid_blockspec.md 与 docs/pallas/quickstart.md;Pallas 的总体设计见 docs/pallas/design.md。动手练习建议:实现一个把其他维度也流水线化的sum内核,并为add和sum内核补充 Megacore 注解。
- 机器学习
- 深度学习
【免费下载链接】jax
Composable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more
相关推荐
JAX Pallas TPU 快速入门:内存空间、Ref 内核与流水线并行(emit_pipeline 实战)
JAX Pallas TPU 快速入门:内存空间、Ref 内核与流水线并行(emit_pipeline 实战) 本篇指南以 JAX 开源仓库中的 docs/pa
人工智能机器学习深度学习编译器高性能计算JAX Pallas TPU 流水线编程实战:内存层级、多重缓冲、动态块形状与 Megacore 调度
JAX Pallas TPU 流水线编程实战:内存层级、多重缓冲、动态块形状与 Megacore 调度 本文围绕 JAX 仓库中 TPU 专属的 Pallas
人工智能机器学习深度学习编译器高性能计算JAX Pallas SparseCore 内核编写指南:在 TPU 稀疏核上实现 gather、scatter 与流水线内核
JAX Pallas SparseCore 内核编写指南:在 TPU 稀疏核上实现 gather、scatter 与流水线内核 SparseCore(稀疏核)是
人工智能机器学习深度学习编译器高性能计算
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考