【Bug已解决】Kandinsky5 pipeline does not support device_map=balanced 解决方案
一、现象长什么样
想用 accelerate 的device_map="balanced"把 Kandinsky5(一个多组件文生图模型)自动切到多张卡,但加载时失败:
from diffusers import Kandinsky5Pipeline pipe = Kandinsky5Pipeline.from_pretrained( "some/kandinsky5", device_map="balanced", # 期望自动把组件分到多卡 )报错:
ValueError: Kandinsky5Pipeline does not support device_map. Please load with `.to(device)` instead.或者:
RuntimeError: component 'prior' and 'decoder' ended up on different devices; generate() mismatched也可能:加载看似成功,但pipe("a cat")时各组件(prior / decoder / text_encoder / image_encoder / movq)散在不同卡,forward 里张量跨设备,RuntimeError: tensors on different devices。
最迷惑的是:同家族的 SDXL/Flux 能device_map="balanced",偏偏 Kandinsky5 不行——因为它的组件比一般 pipeline 多(prior + decoder 是两个独立子网络,还有 image_encoder),pipeline 的from_pretrained没把device_map正确透传到每个组件,或组件间没做设备对齐。
二、背景
device_map="balanced"(accelerate)的语义是:根据各子模块参数量,把模型自动切分到多张卡(或 CPU 卸载),让每张卡负载均衡。这需要:
- pipeline 支持 device_map:
from_pretrained把device_map透传给每个子组件(prior/decoder/text_encoder/...)的from_pretrained。 - 组件可被切分:每个子组件本身能被 accelerate 切分(即它是标准
nn.Module,有清晰的子模块边界)。 - 运行时设备对齐:推理时,输入张量被送到「第一个被用的组件所在设备」,且组件之间若被切到不同卡,accelerate 的 hook 会处理跨设备传输——但这要求 pipeline 的
generate/__call__不手动把张量.to()到错误设备。
Kandinsky5 的问题:
- 它的
from_pretrained没处理device_map参数(老代码只认torch_dtype/low_cpu_mem_usage),于是要么直接拒,要么忽略 device_map、组件全在 CPU/单卡,但generate假设多卡而错位。 - 或者组件 prior/decoder 是两个独立 pipeline 类,各自
from_pretrained时 device_map 没统一传,一个上了 GPU、一个留 CPU。
根子是:Kandinsky5 pipeline 的from_pretrained没把device_map透传到多组件、且generate没做设备对齐,导致device_map="balanced"被拒或运行时跨设备崩。
三、根因
根因一句话:Kandinsky5 的多组件结构(prior/decoder/image_encoder/movq)让device_map="balanced"的透传与运行时设备对齐变复杂,pipeline 的from_pretrained没把 device_map 正确传给每个组件、generate也没对齐设备,导致被拒或跨设备崩溃。
三点展开:
- device_map 未透传:
from_pretrained没把device_map传给 prior/decoder 各自的加载,组件设备错乱。 - 多组件难切分:prior 和 decoder 是两套网络,accelerate 切分时边界处理复杂,pipeline 没适配。
- 运行时未对齐:
generate把输入.to()到错误设备,或组件间跨设备传输没靠 accelerate hook。
不是卡不够,是「多组件 device_map 透传 + 对齐」缺失。
四、最小可运行复现
不依赖真实模型,模拟「device_map 未透传导致组件设备错位」:
from dataclasses import dataclass from typing import Optional @dataclass class FakeComponent: device: str = "cpu" @dataclass class FakeKandinsky5: prior: FakeComponent = None decoder: FakeComponent = None _device_map = None @classmethod def from_pretrained(cls, device_map=None): self = cls(prior=FakeComponent(), decoder=FakeComponent()) # 错误:没把 device_map 透传给组件 if device_map is not None: self._device_map = device_map # 假装只 prior 上了 gpu,decoder 留 cpu self.prior.device = "cuda:0" # decoder 忘记处理 -> 仍在 cpu return self def generate(self): # 运行时:prior 在 cuda,decoder 在 cpu -> 跨设备 if self.prior.device != self.decoder.device: raise RuntimeError("组件设备错位: " f"prior={self.prior.device} decoder={self.decoder.device}") return "image" pipe = FakeKandinsky5.from_pretrained(device_map="balanced") try: pipe.generate() except RuntimeError as e: print("device_map 错位炸:", e)跑出来:prior 上 cuda、decoder 留 cpu,generate 跨设备RuntimeError。这就是「多组件 device_map 失败」的精确复现。
五、解决方案(第一层:最小直接修复)
最小修复:Kandinsky5 的from_pretrained把device_map透传给每个子组件加载;generate前用 accelerate 的cpu_offload/device_maphook 保证组件间设备一致,或统一把输入送到首个组件设备。
from diffusers import DiffusionPipeline import torch class Kandinsky5Pipeline(DiffusionPipeline): def __init__(self, prior, decoder, text_encoder, image_encoder, movq): super().__init__() self.register_modules(prior=prior, decoder=decoder, text_encoder=text_encoder, image_encoder=image_encoder, movq=movq) @classmethod def from_pretrained(cls, pretrained_model_name_or_path, device_map=None, **kw): # 关键:device_map 透传给每个子组件 sub_kwargs = dict(kw) if device_map is not None: sub_kwargs["device_map"] = device_map prior = cls.Prior.from_pretrained( f"{pretrained_model_name_or_path}/prior", **sub_kwargs) decoder = cls.Decoder.from_pretrained( f"{pretrained_model_name_or_path}/decoder", **sub_kwargs) # image_encoder / text_encoder 同样透传 return super().from_pretrained(pretrained_model_name_or_path, prior=prior, decoder=decoder, **kw) @torch.no_grad() def __call__(self, prompt, **kw): # 运行时设备对齐:把输入送到 prior 所在设备 dev = next(self.prior.parameters()).device # 各组件已在 from_pretrained 经 device_map 切好,accelerate hook 处理跨设备 prior_out = self.prior(self._encode(prompt).to(dev)) dec_out = self.decoder(prior_out.to(next(self.decoder.parameters()).device)) return dec_out要点:
device_map透传给 prior/decoder/image_encoder 各自的from_pretrained,组件按 accelerate 切分。generate不手动把张量乱.to(),交给 accelerate 的 device_map hook 处理跨设备。- 若 accelerate 不支持某组件切分,则退化为「统一
.to(主设备)」并告警。
这一步单独就让device_map="balanced"在 Kandinsky5 上可用。
六、解决方案(第二层:结构性改进)
第一层是「改一个 pipeline 的 from_pretrained」。但 diffusers 多组件 pipeline(Kandinsky 系列、Stable Cascade 等)都面临同样问题。更稳的做法把「多组件 device_map 透传 + 设备对齐」收敛成单一策略。
from dataclasses import dataclass, field from typing import Dict, List, Optional @dataclass class KandinskyDeviceMapPolicy: """多组件 pipeline device_map 透传与对齐的单一策略。""" # 可切分的组件名 shardable: List[str] = field(default_factory=lambda: [ "prior", "decoder", "image_encoder", "text_encoder", "movq", ]) def submodule_kwargs(self, device_map: Optional[str], base: dict) -> dict: if device_map is None: return base return {**base, "device_map": device_map} def align_devices(self, components: Dict[str, "torch.nn.Module"]) -> Dict[str, str]: """返回每个组件实际设备;校验是否需统一。""" devs = {} for name, mod in components.items(): if name in self.shardable and mod is not None: devs[name] = next(mod.parameters()).device.type return devs def assert_runnable(self, components: Dict[str, "torch.nn.Module"]): devs = self.align_devices(components) unique = set(devs.values()) if len(unique) > 1: # 多设备:依赖 accelerate hook;若不允许则统一主设备 main = next(iter(devs.values())) return False, f"组件分布在 {unique},需 accelerate hook 或统一到 {main}" return True, "设备一致" # 用法 policy = KandinskyDeviceMapPolicy() sub_kw = policy.submodule_kwargs("balanced", {"torch_dtype": "auto"}) # 各组件 from_pretrained(..., **sub_kw) ok, msg = policy.assert_runnable({"prior": prior, "decoder": decoder})结构收益:
- 单一策略:device_map 透传、组件设备对齐集中在
KandinskyDeviceMapPolicy。 - 可校验:
assert_runnable生成前断言设备一致性,避免运行时跨设备崩。 - 可扩展:新多组件 pipeline 复用,shardable 列表按模型调。
七、解决方案(第三层:断言 / CI 守护)
写 pytest 守三条:(1) device_map 透传到子组件;(2) 多设备时依赖 hook/统一;(3) 单设备一致可运行。
import pytest from your_lib import KandinskyDeviceMapPolicy def test_device_map_passed_to_submodules(): p = KandinskyDeviceMapPolicy() kw = p.submodule_kwargs("balanced", {"torch_dtype": "auto"}) assert kw["device_map"] == "balanced" assert kw["torch_dtype"] == "auto" def test_no_device_map_returns_base(): p = KandinskyDeviceMapPolicy() kw = p.submodule_kwargs(None, {"torch_dtype": "auto"}) assert "device_map" not in kw def test_multi_device_detected(): p = KandinskyDeviceMapPolicy() class M: def __init__(self, d): self._d = d def parameters(self): class P: def __init__(self, d): self.device = torch.device(d) yield P(self._d) ok, msg = p.assert_runnable({"prior": M("cuda"), "decoder": M("cpu")}) assert ok is False assert "cuda" in msg and "cpu" in msg def test_single_device_ok(): import torch p = KandinskyDeviceMapPolicy() class M: def parameters(self): class P: device = torch.device("cuda") yield P() ok, msg = p.assert_runnable({"prior": M(), "decoder": M()}) assert ok is TrueCI 常驻跑这四条后,任何「device_map 又没透传」「多设备未对齐」的回归都会立刻爆红。
八、排查清单
Kandinsky5device_map="balanced"失败时按顺序查:
- 先确认报错是
does not support device_map或tensors on different devices——定位透传/对齐。 - 确认
from_pretrained把device_map透传给 prior/decoder/image_encoder 各自加载。 generate不手动把张量乱.to(),交给 accelerate device_map hook 处理跨设备。- 若某组件不可切分,退化为统一
.to(主设备)并告警,而非硬上 device_map。 - 用
assert_runnable生成前校验组件设备一致性。 - 多组件 pipeline(Kandinsky/Stable Cascade)共用
KandinskyDeviceMapPolicy。 - 升级 diffusers/accelerate 后,跑「device_map=balanced 多卡加载 + 生成」冒烟。
九、小结
Kandinsky5 不支持device_map="balanced"根子是多组件(prior/decoder/image_encoder/movq)结构让 device_map 透传与运行时设备对齐复杂化,pipeline 的from_pretrained没把 device_map 传给每个组件、generate也没对齐设备,导致被拒或跨设备崩。修复三层次:第一层把 device_map 透传各子组件、generate 交给 accelerate hook 处理跨设备;第二层用KandinskyDeviceMapPolicydataclass 把多组件 device_map 透传与对齐收敛为单一策略;第三层用 pytest 守「透传」「多设备检测」「单设备可运行」。
工程启示:任何「多组件 pipeline」接device_map自动切分,都必须把 device_map 透传到每个子组件的加载,并让运行时依赖 accelerate 的跨设备 hook 而非手动.to()。组件越多,透传与对齐越容易漏——做成单一策略 + 生成前设备校验,是这类 pipeline 支持多卡/卸载的关键。