Skip to content

Keep the torch bounds inside lerobot's range - #19

Merged
Vertax42 merged 1 commit into
mainfrom
fix/torch-pin-lerobot-compat
Sep 13, 2026
Merged

Vertax42 merged 1 commit into
mainfrom
fix/torch-pin-lerobot-compat

Conversation

@Vertax42

Copy link
Copy Markdown
Collaborator

问题

uv pip install -e . 在 main 上已经无法求解:

Because lerobot>=0.5.0,<=0.5.1 depends on torchvision>=0.21.0,<0.26.0
and openpi==0.1.0 depends on lerobot>=0.5,<0.6, we can conclude that
openpi==0.1.0 depends on torchvision>=0.21.0,<0.26.0.
And because openpi==0.1.0 depends on torchvision>=0.26,<0.27, we can
conclude that openpi==0.1.0 cannot be used.

三个约束对 lerobot 全是空集,uv 只报了第一个撞上的:

包 openpi (#18 后) lerobot 0.5.1 交集
torch >=2.11,<2.12 >=2.2.1,<2.11.0 空
torchvision >=0.26,<0.27 >=0.21.0,<0.26.0 空
torchcodec >=0.11,<0.12 >=0.2.1,<0.11.0 空

#18 之前的上界正好就是 lerobot 的天花板。#18 把三个全推过去了,但没动 lerobot>=0.5,<0.6。

CI 发现不了:test.yml 只装 omegaconf 然后解析 YAML,从不求解 openpi 的依赖。

为什么 9.19 不是必需的

#18 升 torch 不是为了 torch,而是为了拿到 cuDNN 9.19 —— torch 精确 pin nvidia-cudnn-cu12(2.10 → ==9.10.2.21,2.11 → ==9.19.0.56),而 jax-cuda12-plugin 只要 >=9.1,<10,所以 torch wheel 决定了 JAX 加载哪个 cuDNN。pyproject 的原注释也承认 torch 在这里只跑 DataLoader。

逐条核对那个版本的必要性:

  • 真实门槛是 9.5。 jax/_src/cudnn/fused_attention_stablehlo.py 里 H_max = 256 if cudnn_version >= 90500 and is_on_hopper else 128,pi0.5 是 head_dim=256。torch 2.10 带的 9.10.2 早就过线。
  • 9.14 的 NaN q-gradient 已在代码层解决 —— gemma._stop_gradient_for_fully_masked_queries 就是干这个的,它的 docstring 写明「fixes that without touching attn_mask」。
  • bf16 发散与 cuDNN 版本无关 —— 同一个 docstring 说「Neither mask variant cures the 2026-08-30 bfloat16 divergence」,解法是 float16 custom VJP。

真正出过事的是多来源混装(9.10.2 dispatcher 套 9.14 engine,版本号还查不出来,仍报 91400),那是 check_cuda_stack.py section 2 的职责,和版本号高低无关。

改动

  • pyproject.toml:三个上界回到 lerobot 的天花板;注释改成说明门槛出自 JAX 而非 torch
  • scripts/check_cuda_stack.py:EXPECTED_CUDNN = (9, 19) 精确 pin → MIN_CUDNN = 90500 下界断言,注释指明出处;混合来源检查(section 2)和生产 shape 前后向(section 3/4)原样保留
  • README.md / docs/training-optimization.md:同步

验证

uv pip install -e . --dry-run  →  Resolved 185 packages in 2.93s,torch 相关零变动
pytest (config + policies + examples)  →  131 passed

现有环境 torch 2.10.0+cu128 + cuDNN 9.10.2.21(91002)同时满足新约束和 ≥ 90500 门槛,不需要重装任何 torch 包。

门槛逻辑单测:9.5.0 / 9.10.2 / 9.14 / 9.19 通过,9.4 / 9.1 拒绝。

唯一失败的 policy_rtc_test 在干净树上同样失败,与本改动无关(#17 的描述里也记录了这点)。

未验证

FP16 cuDNN 那条路径的 3,000 步收敛验证当初是在 9.19 上做的,换到 9.10.2 后应当重跑一次再用于生产。use_cudnn_attention 默认 false,现有配置不受影响。

🤖 Generated with Claude Code

`uv pip install -e .` no longer resolves on main:

    Because lerobot>=0.5.0,<=0.5.1 depends on torchvision>=0.21.0,<0.26.0
    and openpi==0.1.0 depends on lerobot>=0.5,<0.6, we can conclude that
    openpi==0.1.0 depends on torchvision>=0.21.0,<0.26.0.
    And because openpi==0.1.0 depends on torchvision>=0.26,<0.27, we can
    conclude that openpi==0.1.0 cannot be used.

All three bounds are empty against lerobot, not just the one uv reports first:
torch <2.11.0 vs >=2.11, torchvision <0.26.0 vs >=0.26, torchcodec <0.11.0 vs
>=0.11. The pre-#18 bounds were exactly lerobot's ceilings; #18 pushed all three
past them without touching `lerobot>=0.5,<0.6`. CI never caught it because
test.yml installs only omegaconf and parses YAML -- it never resolves openpi.

#18 raised torch to reach cuDNN 9.19, not for anything in torch: torch pins
nvidia-cudnn-cu12 exactly (2.10 -> ==9.10.2.21, 2.11 -> ==9.19.0.56) while
jax-cuda12-plugin only asks >=9.1,<10, so the torch wheel decides what JAX loads.
That version is not required:

- The real floor is 9.5. jax/_src/cudnn/fused_attention_stablehlo.py caps
  `H_max` at 128 unless `cudnn_version >= 90500` and the GPU is Hopper; pi0.5
  runs head_dim=256. torch 2.10's 9.10.2 clears it.
- The 9.14 NaN q-gradients are already handled in code -- that is what
  gemma._stop_gradient_for_fully_masked_queries does.
- The bf16 divergence is a kernel-precision problem the float16 custom VJP
  addresses, not a version problem.

What actually broke was a mixed-source stack (a 9.10.2 dispatcher over 9.14
engines, which still reports 91400), and no version number reveals that --
check_cuda_stack.py section 2 is the check for it.

So: bounds back to lerobot's ceilings, and check_cuda_stack.py asserts the
`>= 9.5` floor with its source named instead of an exact 9.19 pin.

Verified: `uv pip install -e . --dry-run` resolves 185 packages with no torch
change (the installed 2.10.0+cu128 / cuDNN 9.10.2.21 already satisfies both the
new bounds and the floor); 131 passed across config, policies and examples.

Not verified: the FP16 cuDNN path's 3,000-step convergence run was done on 9.19.
Re-run it on 9.10.2 before enabling use_cudnn_attention in production. It stays
false by default, so no current config is affected.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
@Vertax42
Vertax42 merged commit 7396c33 into main Sep 13, 2026
1 check passed
@Vertax42

Copy link
Copy Markdown
Collaborator Author

@Hubo1231 已合并(main 当时装不上,先解除 breakage)。

这个改动把 torch/torchvision/torchcodec 的上界退回 lerobot 的天花板,依据是 cuDNN 的真实硬门槛只有 9.5 —— jax/_src/cudnn/fused_attention_stablehlo.py 里 H_max = 256 if cudnn_version >= 90500 and is_on_hopper else 128,pi0.5 是 head_dim=256,torch 2.10 带的 9.10.2 已经过线。

另外两条也核对过:9.14 的 NaN q-gradient 由 _stop_gradient_for_fully_masked_queries 在代码层解决,bf16 发散由 float16 custom VJP 解决,都不依赖具体 cuDNN 版本。

但这些都是从代码里读出来的。如果你在 9.10.2 上实际撞过什么问题、9.19 是有意选的,说一声,我们再调整 —— 比如改成直接依赖 nvidia-cudnn-cu12==9.19.* 而不动 torch,这样既拿到 9.19 又不越过 lerobot 的上界。

还有一件事要麻烦你:FP16 那条路径的 3000 步收敛验证当初是在 9.19 上做的,现在环境是 9.10.2。启用 use_cudnn_attention 之前建议在 H100 上重跑一次。默认值仍是 false,现有配置不受影响。

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant