diff --git a/comfy/ldm/minimax/model.py b/comfy/ldm/minimax/model.py index b6feb8860..14baed3bd 100644 --- a/comfy/ldm/minimax/model.py +++ b/comfy/ldm/minimax/model.py @@ -27,6 +27,8 @@ import comfy.patcher_extension import comfy.quant_ops from comfy.ldm.modules.attention import AttentionTensorContainer, optimized_attention +_AMD_ARCH_CACHE = {} + FRAME_PER_TOKEN = (1, 4, 4, 4, 4) FRAME_RESCALE = 5.0 / 3.0 VISUAL_COND_TIMESTEP = 0.999 @@ -178,13 +180,43 @@ class Attention(nn.Module): q = self.q_norm(q.view(s, self.heads, self.head_dim)) k = self.k_norm(k.view(s, self.heads, self.head_dim)) v = v.clone() - q = AttentionTensorContainer(q.transpose(0, 1).unsqueeze(0)) - k = AttentionTensorContainer(k.transpose(0, 1).unsqueeze(0)) - v = AttentionTensorContainer(v.transpose(0, 1).unsqueeze(0)) + q = q.transpose(0, 1).unsqueeze(0) + k = k.transpose(0, 1).unsqueeze(0) + v = v.transpose(0, 1).unsqueeze(0) + q, k, v = _contiguous_qkv_for_gfx1151(q, k, v) + q = AttentionTensorContainer(q) + k = AttentionTensorContainer(k) + v = AttentionTensorContainer(v) out = optimized_attention(q, k, v, self.heads, mask=None, skip_reshape=True, transformer_options=transformer_options) return self.out_proj(out.squeeze(0)) +def _contiguous_qkv_for_gfx1151(q, k, v): + if not comfy.model_management.is_amd() or q.shape[-2] < 5000 or q.shape[-1] != 128: + return q, k, v + if all(x.is_contiguous() for x in (q, k, v)): + return q, k, v + if _amd_arch(q.device) != "gfx1151": + return q, k, v + return tuple(x.contiguous() for x in (q, k, v)) + + +def _amd_arch(device): + if device.type != "cuda": + return None + index = device.index if device.index is not None else torch.cuda.current_device() + key = (device.type, index) + if key in _AMD_ARCH_CACHE: + return _AMD_ARCH_CACHE[key] + if index < 0 or index >= torch.cuda.device_count(): + _AMD_ARCH_CACHE[key] = None + return None + arch_name = torch.cuda.get_device_properties(index).gcnArchName + arch = arch_name.split(":", 1)[0] + _AMD_ARCH_CACHE[key] = arch + return arch + + class MLP(nn.Module): def __init__(self, hidden, ffn, dtype=None, device=None, operations=None): super().__init__() diff --git a/tests-unit/comfy_test/test_minimax_h3_attention.py b/tests-unit/comfy_test/test_minimax_h3_attention.py new file mode 100644 index 000000000..a684a030a --- /dev/null +++ b/tests-unit/comfy_test/test_minimax_h3_attention.py @@ -0,0 +1,123 @@ +import torch + +import comfy.ldm.minimax.model as minimax_model +import comfy.ops + + +def test_h3_attention_makes_large_qkv_contiguous_only_on_gfx1151(monkeypatch): + attention = minimax_model.Attention( + hidden=8, + heads=2, + head_dim=128, + eps=1e-5, + dtype=torch.float32, + device="cpu", + operations=comfy.ops.disable_weight_init, + ) + x = torch.randn(5000, 8) + layouts = [] + + def capture_attention(q, k, v, heads, **kwargs): + q, k, v = (x.peek() for x in (q, k, v)) + layouts.append(tuple(t.is_contiguous() for t in (q, k, v))) + return q.transpose(1, 2).reshape(1, q.shape[2], heads * q.shape[3]) + + monkeypatch.setattr(minimax_model, "optimized_attention", capture_attention) + monkeypatch.setattr(minimax_model.comfy.model_management, "is_amd", lambda: True) + + arch = ["gfx1100"] + monkeypatch.setattr(minimax_model, "_amd_arch", lambda device: arch[0]) + gfx1100_output = attention(x) + + arch[0] = "gfx1151" + gfx1151_output = attention(x) + + assert layouts == [(False, False, False), (True, True, True)] + torch.testing.assert_close(gfx1151_output, gfx1100_output) + + +def make_qkv(seq=5000, dim=128, contiguous=False): + if contiguous: + return tuple(torch.randn(1, 2, seq, dim) for _ in range(3)) + fused = torch.randn(seq, 6 * dim) + return tuple(x.view(seq, 2, dim).transpose(0, 1).unsqueeze(0) for x in fused.split(2 * dim, dim=-1)) + + +def test_h3_qkv_contiguous_gate_rejects_unvalidated_inputs(monkeypatch): + cases = [ + (False, make_qkv()), + (True, make_qkv(seq=4999)), + (True, make_qkv(dim=64)), + (True, make_qkv(contiguous=True)), + ] + + for is_amd, inputs in cases: + arch_calls = [] + monkeypatch.setattr(minimax_model.comfy.model_management, "is_amd", lambda: is_amd) + monkeypatch.setattr(minimax_model, "_amd_arch", lambda device: arch_calls.append(device)) + outputs = minimax_model._contiguous_qkv_for_gfx1151(*inputs) + assert arch_calls == [] + assert all(output is original for output, original in zip(outputs, inputs)) + + +def test_h3_qkv_contiguous_gate_accepts_5000_boundary(monkeypatch): + monkeypatch.setattr(minimax_model.comfy.model_management, "is_amd", lambda: True) + arch_calls = [] + monkeypatch.setattr(minimax_model, "_amd_arch", lambda device: arch_calls.append(device) or "gfx1151") + inputs = make_qkv(seq=5000) + + outputs = minimax_model._contiguous_qkv_for_gfx1151(*inputs) + + assert arch_calls == [inputs[0].device] + assert all(output is not original for output, original in zip(outputs, inputs)) + assert all(output.is_contiguous() for output in outputs) + + +def test_h3_amd_arch_is_device_aware_cached(monkeypatch): + minimax_model._AMD_ARCH_CACHE.clear() + calls = [] + monkeypatch.setattr(minimax_model.torch.cuda, "device_count", lambda: 2) + + def get_properties(device): + calls.append(device) + return type("Props", (), {"gcnArchName": f"gfx115{device + 1}:sramecc+:xnack-"})() + + monkeypatch.setattr(minimax_model.torch.cuda, "get_device_properties", get_properties) + + assert minimax_model._amd_arch(torch.device("cpu")) is None + assert calls == [] + assert minimax_model._amd_arch(torch.device("cuda:0")) == "gfx1151" + assert minimax_model._amd_arch(torch.device("cuda:0")) == "gfx1151" + assert minimax_model._amd_arch(torch.device("cuda:1")) == "gfx1152" + assert minimax_model._amd_arch(torch.device("cuda:1")) == "gfx1152" + assert calls == [0, 1] + + minimax_model._AMD_ARCH_CACHE.clear() + + +def test_h3_amd_arch_caches_invalid_device_miss(monkeypatch): + minimax_model._AMD_ARCH_CACHE.clear() + calls = [] + monkeypatch.setattr(minimax_model.torch.cuda, "device_count", lambda: 1) + monkeypatch.setattr(minimax_model.torch.cuda, "get_device_properties", lambda device: calls.append(device)) + + assert minimax_model._amd_arch(torch.device("cuda:1")) is None + assert minimax_model._amd_arch(torch.device("cuda:1")) is None + assert calls == [] + assert minimax_model._AMD_ARCH_CACHE == {("cuda", 1): None} + + minimax_model._AMD_ARCH_CACHE.clear() + + +def test_h3_amd_arch_propagates_property_errors(monkeypatch): + minimax_model._AMD_ARCH_CACHE.clear() + monkeypatch.setattr(minimax_model.torch.cuda, "device_count", lambda: 1) + monkeypatch.setattr(minimax_model.torch.cuda, "get_device_properties", lambda device: (_ for _ in ()).throw(RuntimeError("probe failed"))) + + try: + minimax_model._amd_arch(torch.device("cuda:0")) + except RuntimeError as error: + assert str(error) == "probe failed" + else: + raise AssertionError("CUDA property error was silently ignored") + assert minimax_model._AMD_ARCH_CACHE == {}