使用Einops的rearrange替代torch.chunk为何结果不一致?
问题原因及解决方法
你遇到的核心问题是PyTorch的chunk和你用的Einops rearrange拆分逻辑完全不同:
- 原代码里的
th.chunk(emb, 2, dim=1)是在第1维度上做连续拆分:假设该维度长度是C,第一块取前C/2个连续元素,第二块取后C/2个连续元素。 - 你写的
rearrange(emb, "b (split c) ... -> b split c ...", split=2)是做交错拆分:它会把第1维度的元素按位置0,2,4...和1,3,5...分成两组,相当于把原维度的元素“穿插”着拆分,和chunk的连续拆分逻辑完全相反,所以结果自然不一样。
正确的Einops实现方式
如果你想用Einops替代chunk的逻辑,直接用Einops的split函数最贴合原行为,它和PyTorch的chunk逻辑一致:
from einops import split emb = rearrange(emb, "b (c h w) -> b c h w", w=1, h=1) cond_w, cond_b = split(emb, split=2, dim=1)
内容的提问来源于stack exchange,提问作者Foobar
相关产品推荐
相关产品推荐

