Skip to content

[Bug] --cache 与 --parallel ulysses 组合崩溃:_get_Fn_residual 中 hidden_states 维度不匹配(分片 vs 未分片) #1094

Description

@brandneway

现象概述

--cache--parallel ulysses 组合使用时,在首个去噪/warmup 步就会因 cache 残差计算中的张量维度不匹配而崩溃。
同一命令去掉 --cache 则完全正常,所以问题出在 cache 与 ulysses 的交互上。

环境

  • cache_dit:1.5.0
  • torch 2.10.0+cpu、torch_npu 2.10.0、diffusers 0.39.0、accelerate 0.34.2
  • 硬件:4× Ascend 910B4(NPU),单机
  • 模型:Wan2.2-T2V-A14B-Diffusers(bf16,--cpu-offload

复现步骤

正常(不加 cache):

torchrun --nproc_per_node=4 -m cache_dit.generate wan2.2_t2v \
  --model-path /data/Wan2.2-T2V-A14B-Diffusers --cpu-offload --parallel ulysses

报错(加上 --cache):

torchrun --nproc_per_node=4 -m cache_dit.generate wan2.2_t2v \
  --model-path /data/Wan2.2-T2V-A14B-Diffusers --cpu-offload --parallel ulysses --cache

报错信息(4 个 rank 完全相同,已精简到关键日志)

File ".../cache_dit/_utils/registers.py", line 555, in run
    _ = pipe(**warmup_kwargs)
File ".../cache_dit/caching/cache_blocks/pattern_base.py", line 233, in forward
    Fn_hidden_states_residual = self._get_Fn_residual(original_hidden_states, hidden_states)
File ".../cache_dit/caching/cache_blocks/pattern_base.py", line 195, in _get_Fn_residual
    Fn_hidden_states_residual = hidden_states - original_hidden_states.to(hidden_states.device)
RuntimeError: The size of tensor a (5070) must match the size of tensor b (20280) at non-singleton dimension 1

崩溃前还会出现一条 warning:

[Cache-DiT] Expected input tensor to have 2 dimensions, but got 1 dimensions; split will be skipped.

根因分析

20280 / 5070 == 4 == ulysses_size。ulysses 会在 transformer block(Fn blocks)内部沿 head 维对 hidden_states 做切分,因此:

  • original_hidden_states(block 入口处捕获,切分之前)→ 第 1 维 = 20280
  • hidden_states(经过 Fn blocks 之后)→ 第 1 维 = 5070(被 ulysses_size 切分)

_get_Fn_residual 已经对“形状不一致”做了容错,但只针对 context parallelismpattern_base.py):

if self._check_if_context_parallel_enabled(...) and (original_hidden_states.shape != hidden_states.shape):
    Fn_hidden_states_residual = hidden_states          # CP:容忍形状不一致
else:
    Fn_hidden_states_residual = hidden_states - original_hidden_states   # ulysses:在这里崩

_check_if_context_parallel_enabled 仅当模块挂了 ContextParallelSplitHook 时才返回 True;而 ulysses 走的是另一条路径(distributed/config.py 里的 ulysses_size / usp_enabled()),所以对 ulysses 该守卫为 False,直接进入减法分支,两个形状不一致的张量相减导致报错。
已确认该“仅 CP”守卫在 main 分支仍然存在(pattern_base.py:184/195);同样的 hidden_states - original_hidden_states 写法在 414、619 行也有,可能存在同样问题。

不加 --cache 时根本不会调用 _get_Fn_residual,这就是 ulysses 单独能跑通的原因。
该不匹配本质上只是 ulysses 的 head 维切分所致,因此与硬件后端无关(本次在 NPU 上复现)。

期望行为

--cache 应能兼容 --parallel ulysses——例如把 _get_Fn_residual 的形状不一致容错扩展到张量并行(ulysses / usp_enabled()),或在 ulysses 切分之后再捕获 original_hidden_states,使两个张量都处于已切分的形状。

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions