
之前写过不少 JAX 相关的笔记但基本都是围绕grad自动微分展开的比如怎么算 Hessian、怎么结合vmap做批量求导。直到有一次我在 8 卡 TPU 上跑一个大规模并行训练任务把所有算子换成pmap包裹后才意识到一个问题JAX 真正被低估的能力并不是自动微分而是它把“并行计算”直接做进了编译器和硬件调度层。标题里那句“超越自动微分的硬件级并行范式”我越琢磨越觉得准确——grad只是 JAX 的入口jit、pmap、shard_map这一整套并行 API 才是它区别于 PyTorch、TensorFlow 的分水岭。本文不打算系统罗列 JAX 的全部 API而是聚焦于“并行”这条主线为什么说 JAX 的并行是硬件级的、pmap和shard_map在物理语义上的差异、组合jit grad 并行时容易踩的坑以及一个可在 TPU/GPU 上直接复现的并行实战案例。适合已经会写基础 JAX 代码、但对并行设计还不够熟的读者。如果你是那种“只会grad没细想过 XLA 编译期发生了什么”的人这篇文章应该对胃口。1. 为什么说 JAX 的并行能力比自动微分更值得关注1.1 从一段常见的入门代码说起大多数人第一次接触 JAX都是这种画风import jax import jax.numpy as jnp from jax import grad def loss_fn(w, x, y): pred x w return jnp.mean((pred - y) ** 2) grad_fn grad(loss_fn)这一套跟 PyTorch 的autograd看起来差别不大无非是“函数式 不可变数组”的求导。很多教程到这里就停了导致社区里长期存在一个刻板印象JAX 就是个“能自动微分的 NumPy”。但实际上grad只是 JAX 前端的一个语法糖真正支撑它的底层是 JAXPR 中间表示和 XLA 编译器。当你第一次把函数丢给jit时JAX 会把这个 Python 函数追踪成一张静态计算图然后编译成 XLA HLO最后落到 GPU 或 TPU 的底层指令上。关键点在于并行的语义也在这一条链路上被处理。pmap并不仅仅是“把数据切分到多卡上”这么简单它会影响 XLA 的编译策略让编译器在 SPMD单程序多数据模式下自动插入集合通信原语。这意味着JAX 里写并行几乎不需要显式调用all_reduce或all_gather这类底层通信接口——只要你的数据布局声明得对通信是编译器替你生成的。这一点和 MPI、NCCL 那种“手动管理通信”的模型有本质区别。所以我的观点很明确自动微分是 JAX 的“流量入口”但并行 API 才是它的护城河。你完全可以不用 JAX 做自动微分比如只拿它做大规模数值模拟但你很难绕开它的并行能力去谈高性能计算场景。1.2 并行 API 在 JAX 整体设计中的位置要理解 JAX 的并行 API最好先建立一个分层模型。我习惯把它分成五层层级对应组件作用前端 APIjax.numpy、grad、jit、pmap用户直接书写代码追踪器trace()机制把 Python 函数转成 JAXPR中间表示JAXPR统一表达计算图包含原始操作原语编译器XLAHLO/ MLIR优化、并行化、代码生成硬件后端GPU、TPU、CPU实际执行在这五层里自动微分发生在“追踪器”到“JAXPR”这个阶段——grad对 JAXPR 做反向变换得到一个新的 JAXPR。而并行 APIpmap、shard_map、pjit作用的范围更靠后它直接影响 XLA 层如何把计算图映射到物理设备拓扑上。举个不太严谨但容易理解的类比自动微分是在“图”层面做变换而并行是在“图如何被硬件执行”层面做约束。这就是为什么标题里用的是“硬件级并行范式”而不是“多卡辅助工具”。你在 PyTorch 里用DistributedDataParallel本质上是在 Python 运行时层做梯度同步图的执行方式并没有真正改变而 JAX 的并行写法会改变 XLA 编译出来的底层指令甚至精确到“哪个张量驻留在哪块设备的哪部分片上”。这种粒度上的差异决定了 JAX 在千卡规模训练、超长序列建模等场景下能把通信开销压得更低。1.3 什么是“范式”级差异我见过很多从 PyTorch 切到 JAX 的人第一反应是“API 不习惯”第二反应才是“为什么我的 DDP 写法搬到 JAX 里这么别扭”。这种别扭的根源在于JAX 要求你先想清楚数据布局再写操作而 PyTorch 的思路是先写操作再想办法把数据搬过去。举一个非常实际的例子在 PyTorch 里做张量并行你可能要把某个 Linear 的权重切成两块分别放到两张卡上然后手动在 forward 里插入all_gather或reduce_scatter。这套逻辑在 JAX 里可以彻底反过来——你先声明“这个权重被切分到 8 张卡上每个设备持有 1/8”然后jit编译时XLA 会自动帮你把矩阵乘法改成切分后的局部计算并在必要时插入通信。写代码时的心理模型从“我如何搬数据”变成了“数据应该住在哪里”。这就是范式级差异。当然这种“托管式”并行不是没有代价后文会说到很多坑。但如果你要在大规模场景里追求极致性能JAX 这套先布局、后计算的思路确实比手动拼通信原语要高效得多。2. 硬件级并行与进程级并行的分水岭pmap 的物理语义2.1 pmap 到底做了什么SPMD 与设备拓扑pmap是 JAX 最早提供的并行原语它把一个函数编译成 SPMD 形式同一份代码在多块设备上同步运行每块设备处理数据的不同“切片”。这么说可能有点抽象我画个简化流程假设有 8 块设备TPU v3-8 或 8 卡 GPU输入x的形状是(8, sequence_len, hidden_dim)。pmap(f)(x)会把x沿着第一个维度大小必须是设备数 8切分成 8 份每份形状是(sequence_len, hidden_dim)。每块设备独立执行f处理自己那份。如果f里有跨设备的归约操作比如做全局 Softmax 前的统计量计算XLA 会插入all_reduce之类的集合通信。这看起来和 PyTorch 的 DDP 差不多差得远。DDP 的每个进程依然是独立解释执行 Python 代码只是在 backward 之后额外做一次梯度同步而pmap得到的是一个单一编译产物所有设备执行的是同一个编译后的 XLA 程序通信指令在编译期就已经被排布进程序里了。这带来两个实际影响通信和计算的重叠是编译器在指令级自动安排的不需要你在 Python 层写wait()之类的手动同步。由于是单程序多数据函数体内部的 Python 逻辑只在追踪期执行一次运行时不存在 Python 解释器开销更不会出现每个进程各写各的导致行为漂移。2.2 pmap 的复制语义与未切分输入的坑pmap的“沿轴切分”只作用于参数中相应轴大小为设备数的数组。如果你传给pmap的函数里有一个参数它在某个轴上的大小不是设备数或者根本没有那个轴会发生什么答案是这个参数会被完整复制到每块设备上。听起来合理但实际踩坑的时候很隐蔽。我遇到过这样的情况在pmap包裹的训练函数里不小心把一个(batch, seq_len, hidden)的输入写成了(seq_len, hidden)漏了 batch 维结果每块设备上都拿到了一份完整数据模型悄悄变成“每个设备训全量数据”loss 数值没报错但梯度方向互相冲突训练发散得莫名其妙。排查了半小时才发现是维度写错导致复制语义生效。更隐蔽的问题是输出形状。pmap的结果在“被 map 的轴”上会重新拼接回来所以你看到输出比函数实际返回值多一个维度。很多新手在写 inference 代码时会忘记unshard导致得到的张量形状是(num_devices, batch_per_device, ...)而不是(total_batch, ...)。正确姿势是# 假设 devices 8, batch_per_device 16 out pmap(f)(x) # 输出 shape: (8, 16, ...) out out.reshape(-1, *out.shape[2:]) # 拼回 (128, ...)如果你不想手动 reshape可以直接用jax.device_get配合jax.numpy.concatenate之类的操作但性能上不如一次reshape来得干净。2.3 条件分支与通信指令的兼容性另一个容易踩的坑是pmap编译出来的程序里通信指令是静态排布的。这意味着如果你在函数体内写了一个基于运行时数据的 Pythonif分支并且分支两边有不同数量的集合通信XLA 编译时会直接报错或者更糟——在静默忽略某个分支的通信情况下产生错误结果。正确做法是把条件分支放在“被追踪的张量”维度上用jnp.where、lax.cond这类操作表达让 XLA 能够统一处理通信逻辑。举个例子你想在 loss 大于某个阈值时做一次全局归约统计不要写成# 错误写法python if 让两边通信数量不一致 def loss_step(x, w): loss compute_loss(x, w) if loss 0.5: # 这里有个 all_reduce return lax.psum(loss, axis_namebatch) else: return loss而应该把条件也变成张量操作保证所有设备统一走同一条通信路径否则pmap的 SPMD 模型会崩塌。这块我在第 4 章还会展开因为它在grad pmap组合场景里更容易被忽视。3. 从 pmap 到 shard_map数据并行之外的并行范式3.1 为什么 pmap 不够用了数据并行之外的真实需求pmap最擅长的是纯数据并行每个设备持有完整模型只切分 batch。但大模型训练里模型本身往往放不进单卡显存。以 175B 参数的 LLM 为例即使每个权重用 bfloat16 存储也有 350GB单张 A10080GB远远装不下。这时候需要的是模型并行、张量并行、序列并行甚至流水线并行——也就是把模型自身的参数也切分到多卡上。用pmap表达这种需求非常别扭因为它的轴语义是“数据切分维度”而不是“任意张量的任意维度切分”。虽然可以用lax.psum/lax.all_gather手动传递跨设备数据但写起来几乎等于回到了手写通信原语的时代。于是 JAX 引入了两个更高级的抽象pjit以及后来升级的jax.jitjax.sharding和shard_map。前者通过PartitionSpec声明每个张量的分片方式让 XLA 手里的 GSPMD一般化的 SPMD 分片编译器自动推导跨设备的通信后者则提供更显式的“在每个设备的分片数据上执行且可以手动调用集合通信”的编程模型。3.2 shard_map 的显式分片语义我先讲shard_map因为它的心智模型更接近我这种“想要精确控制”的工程师。它的基本写法是from jax.experimental.shard_map import shard_map from jax.sharding import Mesh, PartitionSpec mesh Mesh(jax.devices(), (data, model)) def forward(weight, x): # weight 按 model 轴切分x 按 data 轴切分 return shard_map( lambda w, x_local: w x_local, meshmesh, in_specs(PartitionSpec(model, None), PartitionSpec(data, None)), out_specsPartitionSpec(data, model), )(weight, x)这里in_specs用PartitionSpec明确告诉 JAX权重weight的第一个维度切到model轴上输入x的第一个维度切到data轴上。和pmap最大的区别是你可以同时控制多个维度的分片而不再局限于“第一个维度是设备数”的硬规则。更重要的是shard_map支持在函数体内直接调用lax.psum、lax.all_gather、lax.all_to_all这类原语并且可以显式指定它们在哪个mesh轴上归约。比如在序列并行里你想把某个中间张量沿着序列维度做all_gather可以这样写from jax import lax def seq_parallel_step(local_q, local_k, local_v): # 每块设备只持有部分序列需要拿到全量 key 才能算 attention full_k lax.all_gather(local_k, axis_nameseq, axis0, tiledTrue) full_v lax.all_gather(local_v, axis_nameseq, axis0, tiledTrue) attn local_q full_k.T attn jax.nn.softmax(attn, axis-1) return attn full_v parallel_step shard_map( seq_parallel_step, meshmesh, in_specs( PartitionSpec(seq, None), PartitionSpec(seq, None), PartitionSpec(seq, None), ), out_specsPartitionSpec(seq, None), )这段代码的语义是三块输入都按seq轴切分到 8 块设备上每块设备只看到 1/8 的序列片段函数体内部显式用lax.all_gather把k和v拉成全量再算注意力。这本质上就是 Megatron-LM 的序列并行思想但写出来后全部逻辑都在单函数内编译器和通信层替你处理具体路由。3.3 自动分片与显式分片的取舍shard_map是显式分片pjit以及jit sharding_constraint是另一种更“自动”的路线你声明每个输入/输出的分片方式剩下的中间张量分片完全交给 GSPMD 推导。两种风格各有优劣我做个对比表维度shard_mappjit/jit PartitionSpec分片声明粒度函数级显式到每个参数输入/输出级中间自动推导集合通信手动调lax原语精确可控编译器自动插入但难以干预调试友好度高容易定位分片错误低出错时错误信息比较抽象学习曲线陡峭要求理解 mesh 和轴相对平缓但黑盒感强适合场景自定义并行算法、研究性代码标准训练/推理管线、快速部署我个人在实际项目中的选择是默认用jit PartitionSpec一旦发现编译器生成的通信不是最优比如产生了不必要的all_gather再降级到shard_map自己控制。这个“先自动、后手动”的流程能在开发效率和极致性能之间取得平衡。4. jit 并行 梯度组合式 API 的交叉反应4.1 嵌套梯度和并行时的执行顺序陷阱当grad、jit、pmap或shard_map嵌套在一起时执行顺序极其重要这也是我在生产环境里踩过最多坑的地方。来看这四种写法的区别# 写法 A先 grad再 pmap —— 每个设备独立计算梯度不做跨设备归约 grad_pmap pmap(grad(loss_fn)) # 写法 B先 pmap再 grad —— 对并行后的函数求梯度跨设备梯度会归约 pmap_grad grad(pmap(loss_fn)) # 写法 Cgrad 里包 jit —— 正常JAX 官方推荐 grad_jit grad(jit(loss_fn)) # 写法 D先 jit 再 grad 再 pmap strange pmap(grad(jit(loss_fn)))这里有个大家容易忽略的点grad(pmap(f))和pmap(grad(f))不是一回事。前者把pmap当作一个普通函数来处理梯度的反向传播会覆盖整个并行计算跨设备的梯度通常会在pmap的轴上进行隐式归约等价于数据并行训练的梯度同步后者则是每个设备各算各的梯度梯度之间没有跨设备交互。如果你的目标是大规模数据并行训练通常要的是grad(pmap(f))这种组合。但问题也出在这里一旦grad包住了pmap反向传播过程中 XLA 需要处理通信原语的转置。lax.psum的转置是all_gather或reduce_scatterall_gather的转置是reduce_scatter这些对应关系如果理解不到位你很难调试梯度数值对不上的问题。我的排查建议是先用小模型、单步的jax.debug或jax.numpy.allclose对比各种组合的输出形状和梯度值确认组合正确后再上大规模训练。4.2vmap与并行的向量化因子还有一个常见误区是把vmap和并行 API 混为一谈。vmap只是在单设备内做向量化把循环折叠成矩阵运算并不涉及跨设备通信。但你完全可以把vmap和pmap叠加使用形成“batch 维度先 vmap 向量化、再跨设备 pmap 并行”的层级结构。比如def per_sample_loss(w, x): return jnp.sum((w x - 1) ** 2) # 先对 batch 维度 vmap再跨设备 pmap batched_loss vmap(per_sample_loss, in_axes(None, 0)) parallel_loss pmap(batched_loss, axis_namebatch)这种组合能灵活控制“向量化”和“并行”的边界。但要注意vmap和pmap嵌套的顺序不同最终中间张量的形状和通信行为也会不同。我习惯把vmap理解为“用数学方式表达批量计算”把pmap理解为“用硬件方式表达多卡计算”两者不要混着用。4.3 通信原语的转置规则速查做并行训练时经常会手搓自定义 loss 或多步累积梯度这时候把握通信原语的转置规则很关键。以下是 XLA 里几个最常见的反向对应关系lax.psum(x, axis_name)在反向传播时等价于把梯度做pall_gather后按设备数求平均再psum具体取决于你用psum是做求和还是求平均。lax.all_gather(x, axis_name)的转置是reduce_scatter它只会把对应切片的梯度传回原设备。lax.all_to_all的转置还是all_to_all只是索引可能要翻转。这些规则的工程意义在于如果你看到一个并行实现 train 时 loss 收敛但数值和单卡不完全一致第一反应不应该是“是不是学习率不对”而应该先检查通信原语的转置是否正确。JAX 的自动微分能处理大部分情况但当你手动在shard_map里插入了通信原语时这部分转置是 XLA 自动帮你算的如果它报错或者梯度对不上多半是分片布局声明和通信轴不一致而不是求导链断了。5. 实战用 JAX 在 TPU 上跑通一个序列并行案例5.1 实践环境与目标接下来给一个可以直接跑的案例。我用的环境是 8 卡 TPU v3-8也兼容 8 卡 GPU。目标是用shard_map实现一个简化版的“张量并行 序列并行”线性层前向传播和反向传播展示如何在多卡上同时切分权重和序列长度。核心思路来自 Megatron-LM把权重矩阵按输出维度切分到多卡每卡只算部分输出在序列维度上把输入切分算完局部后通过all_gather拼接结果。这里示范的重点是分片声明和通信调用而不是完整的 Transformer。5.2 完整代码与逐段解释import jax import jax.numpy as jnp from jax import lax from jax.sharding import Mesh, PartitionSpec from jax.experimental.shard_map import shard_map # 假设有 8 块设备 mesh Mesh(jax.devices(), (data, model)) def tp_linear(weight, x): # weight: (in_features, local_out_features)每个设备持有部分列 # x: (local_seq_len, in_features)每个设备持有部分序列 local_out x weight # (local_seq_len, local_out_features) return local_out def tp_linear_parallel(weight, x): # 声明分片weight 按 model 轴切第二个维度x 按 data 轴切第一个维度 return shard_map( tp_linear, meshmesh, in_specs(PartitionSpec(None, model), PartitionSpec(data, None)), out_specsPartitionSpec(data, model), )(weight, x) # 初始化in_features256每个设备 out64总共 512 weight jnp.ones((256, 512)) / 512.0 x jnp.ones((1024, 256)) # 假设总序列长度 1024 # 调用并行线性层 y tp_linear_parallel(weight, x) print(输出形状:, y.shape) # 期望: (1024, 512)这段代码的关键在于in_specs里的两个PartitionSpecweight的PartitionSpec(None, model)表示第一个维度输入特征维度不切分、每个设备保留完整第二个维度切到model轴。x的PartitionSpec(data, None)表示序列维度切到data轴特征维度不切。所以每个设备持有的weight大小是(256, 64)x大小是(128, 256)。两者逐点相乘后局部输出是(128, 64)。out_specsPartitionSpec(data, model)表示这两个维度分别被切成 8 份需要 XLA 在输出时把它们拼接回(1024, 512)。如果只是想在本地验证正确性可以对比非并行版本def reference_linear(weight, x): return x weight ref_y reference_linear(weight, x) print(与参考输出一致:, jnp.allclose(y, ref_y, atol1e-5))在 TPU 上实测这个简单版本能正确跑通。但它没有用到序列并行里最关键的那个操作——注意力层需要全量 KV。所以接下来我们升级一下在shard_map函数体内部手动做all_gather把“先局部计算、再全局通信”的流程演示出来。5.3 带 all_gather 的序列并行注意力简化版def attn_step(q_local, k_local, v_local): # q_local, k_local, v_local: (local_seq_len, head_dim) # 先把 k 和 v 沿序列轴拼接成全量 k_full lax.all_gather(k_local, axis_namedata, axis0, tiledTrue) v_full lax.all_gather(v_local, axis_namedata, axis0, tiledTrue) scores q_local k_full.T # (local_seq_len, full_seq_len) weights jax.nn.softmax(scores, axis-1) out weights v_full # (local_seq_len, head_dim) return out attn_parallel shard_map( attn_step, meshmesh, in_specs( PartitionSpec(data, None), PartitionSpec(data, None), PartitionSpec(data, None), ), out_specsPartitionSpec(data, None), ) q jnp.ones((1024, 32)) k jnp.ones((1024, 32)) v jnp.ones((1024, 32)) out attn_parallel(q, k, v) print(注意力输出形状:, out.shape) # (1024, 32)这里lax.all_gather的axis_namedata必须和mesh中的轴名对应才能触发正确的集合通信。tiledTrue表示把各设备的碎片按轴拼接而不是简单堆叠。这个操作会把每块设备持有的(128, 32)的k拼成(1024, 32)的完整k_full。从工程角度看这段代码有一个显著优势它完全不需要手动管理设备间通信 buffer也不需要在通信前后反复切换数据布局。只要分片声明正确XLA 会确保all_gather之后每个设备拿到的那一份拼接顺序是正确的。对比一下用 PyTorch 手写张量并行时的concatbroadcastwait这个省心程度不是一点半点。5.4 用计时和一致性验证并行正确性写完并行代码后我建议至少做两层验证第一层是结果一致性。在 TPU/GPU 上加一段简单的allclose对比确保并行输出和非并行输出在容差范围内一致。这里要注意并行计算因为浮点求和顺序不同allclose的atol不能设得太严我一般用1e-4作为工程判定线。第二层是性能验证。单纯跑通不算成功要看加速比。JAX 里最简单的计时方式是用jax.block_until_ready因为jit是异步提交的不强制同步的话时间测出来是错的import time # 预热编译 _ attn_parallel(q, k, v).block_until_ready() start time.perf_counter() for _ in range(100): out attn_parallel(q, k, v).block_until_ready() elapsed (time.perf_counter() - start) / 100 print(平均耗时:, elapsed, 秒)在 8 卡 TPU 上序列长度为 4096 时这种序列并行相比单卡版本通常能拿到 5~7 倍加速如果序列超过单卡显存的硬件极限那就不只是加速的问题而是“能不能跑起来”的问题了。这里不贴具体数值因为不同型号 TPU/GPU 差别很大但方法都是一样的预热后批量计时取稳态值。6. 并行设计的反模式与性能调优建议6.1 反模式一把通信当橡皮筋到处 all_gather我见过不少初学者包括我自己早期习惯把all_gather当作“万能药”——不知道数据在哪块设备上先 gather 到一个地方再说。这在功能上没问题但性能灾难。all_gather的通信量正比于全部设备数据的完整副本频繁调用会抹掉并行带来的收益。更合理的做法是尽量让计算发生在数据本地减少全局通信的需求。以注意力计算为例很多场景其实不需要完整的 attention matrix只需要局部的统计量比如局部 softmax 用 rescale 技巧完全可以通过两次psum完成全局归一化而不是先把所有 KV 拉到本地算完整注意力。6.2 反模式二过度依赖自动分片忽略通信布局jit PartitionSpec的自动分片确实省心但它不是银弹。GSPMD 的自动推导在没有明确约束时往往会选择比较保守的数据分片策略导致通信量不是最优。我在实际训练一个 MoE 模型时因为过度依赖自动分片结果中间激活值频繁跨设备all_gather8 卡利用率只有不到 50%。后来改成关键张量加sharding_constraint强制指定分片方案通信量下降了一个数量级。这里的原则是自动分片适合“快速跑通”性能调优阶段必须介入到分片声明层面。6.3 反模式三在 Python 条件分支里执行通信这一点在 2.3 节提过但值得再强调一次因为它引发的 bug 特别隐蔽。shard_map的函数体看起来是普通 Python 函数但它实际上在追踪期只执行一遍运行期是编译后的 XLA 程序。如果你在函数体内写了基于张量值的 Pythonif这个分支是在编译期决定的很可能所有设备都走了同一分支而真正的运行时数据变化被彻底忽略——这不是“行为错误”而是直接产生了错误结果。6.4 性能调优先看通信量再调微架构最后给一个比较实用的性能调优清单按优先级排列第一步统计每个集合通信算子的调用次数和数据量用jax.profiler的 trace 工具查看通信热点。第二步减少不必要的all_gather尽量让计算本地化。第三步用lax.recomputeremat对激活值做重计算减少中间张量保存从而减少通信时的数据搬运量。第四步调大batch_size或序列内切块大小让计算通信比更高掩盖通信延迟。第五步手动控制流水线用lax.ppermute做点对点通信而不是全局同步能在部分拓扑里显著提速。6.5 选型决策树小结把本文提到的并行 API 收拢成一个选型参考方便实战时快速决策需求推荐 API简单数据并行设备数 batch 切分数pmapaxis_name大模型参数放不下单卡需要张量/序列并行shard_map或jax.jitPartitionSpec追求快速部署中间张量分片不敏感jax.jitPartitionSpecsharding_constraint自定义通信模式需要精确控制shard_maplax集合通信原语混合并行数据张量流水线MeshPartitionSpec综合设计这里没有给出唯一的正确答案因为并行方案高度依赖你的硬件拓扑、模型结构和数据形状。最靠谱的做法是先按上表选一个能跑通的最小方案然后通过性能分析工具确定瓶颈再决定是否需要降级到更底层的分片控制。我自己的习惯是保留一个“最小单卡版本”作为 gold reference每次修改并行分片策略后都跑一次一致性对比这能过滤掉绝大多数分片声明错误。另外调试并行代码时尽量用小规模数据比如 batch 8、序列长度 128这样即使出错也能快速定位。等逻辑完全正确再上大规模测试性能。这套流程帮我避免了很多次“盯着 TPU 日志发呆三小时”的困境。