perf(minimax): make guarded H3 QKV contiguous on gfx1151

This commit is contained in:
allenliang2022 2026-08-12 00:58:04 +08:00
parent 62b3c94bd4
commit 04f344546b
2 changed files with 158 additions and 3 deletions

View File

@ -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
@ -166,13 +168,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__()

View File

@ -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 == {}