Uploaded sanitized BC250/ROCm Repository.

This commit is contained in:
Fabian
2026-08-20 00:45:43 +02:00
parent 7d2184f1e8
commit d7d22e93b3
678 changed files with 65963 additions and 1 deletions
@@ -0,0 +1,121 @@
"""Test GPU operations that occur between Linear layers in the Lumina2 model.
This tests LayerNorm, RMSNorm, SiLU, tensor reshaping, and manual SDPA."""
import torch
import torch.nn.functional as F
import time, os, sys
os.environ['HSA_OVERRIDE_GFX_VERSION'] = '10.1.0'
os.environ['HSA_ENABLE_SDMA'] = '0'
os.environ['TORCHDYNAMO_DISABLE'] = '1'
print(f"Device: {torch.cuda.get_device_name(0)}", flush=True)
print(f"AMD_SERIALIZE_KERNEL={os.environ.get('AMD_SERIALIZE_KERNEL','unset')}", flush=True)
# Dimensions matching Lumina2 z_image_turbo
B, SEQ, DIM = 1, 1024, 3840
HEADS = 24
HEAD_DIM = DIM // HEADS # 160
dtype = torch.float16
def test_op(name, fn):
print(f"\n--- {name} ---", flush=True)
try:
t0 = time.time()
result = fn()
torch.cuda.synchronize()
dt = time.time() - t0
if hasattr(result, 'shape'):
print(f" OK: {dt*1000:.1f}ms, shape={list(result.shape)}, dtype={result.dtype}", flush=True)
else:
print(f" OK: {dt*1000:.1f}ms", flush=True)
return result
except Exception as e:
print(f" ERROR: {type(e).__name__}: {e}", flush=True)
return None
# Create typical model tensors
x = torch.randn(B, SEQ, DIM, dtype=dtype, device='cuda:0')
torch.cuda.synchronize()
print(f"Input tensor on GPU: {list(x.shape)}", flush=True)
# Test 1: F.layer_norm
weight_ln = torch.randn(DIM, dtype=dtype, device='cuda:0')
bias_ln = torch.randn(DIM, dtype=dtype, device='cuda:0')
test_op("F.layer_norm", lambda: F.layer_norm(x, (DIM,), weight_ln, bias_ln))
# Test 2: RMSNorm manual (common in Lumina2)
def rms_norm(x, w, eps=1e-5):
rms = torch.sqrt(torch.mean(x * x, dim=-1, keepdim=True) + eps)
return x / rms * w
test_op("RMSNorm manual", lambda: rms_norm(x, weight_ln))
# Test 3: SiLU activation
test_op("F.silu", lambda: F.silu(x))
# Test 4: GELU activation
test_op("F.gelu", lambda: F.gelu(x))
# Test 5: Tensor reshape (view + transpose)
def test_reshape():
q = x.view(B, SEQ, HEADS, HEAD_DIM).transpose(1, 2) # [B, HEADS, SEQ, HEAD_DIM]
return q
test_op("view+transpose", test_reshape)
# Test 6: Tensor repeat (GQA expansion)
def test_repeat():
kv = torch.randn(B, SEQ, 8, HEAD_DIM, dtype=dtype, device='cuda:0')
kv = kv.unsqueeze(3).repeat(1, 1, 1, 3, 1).flatten(2, 3)
return kv
test_op("unsqueeze+repeat+flatten", test_repeat)
# Test 7: torch.cat
def test_cat():
a = torch.randn(B, SEQ, DIM, dtype=dtype, device='cuda:0')
b = torch.randn(B, SEQ, DIM, dtype=dtype, device='cuda:0')
return torch.cat([a, b], dim=1)
test_op("torch.cat", test_cat)
# Test 8: Full manual SDPA (like our patch does)
def test_manual_sdpa():
Q = torch.randn(B, HEADS, SEQ, HEAD_DIM, dtype=dtype, device='cuda:0')
K = torch.randn(B, HEADS, SEQ, HEAD_DIM, dtype=dtype, device='cuda:0')
V = torch.randn(B, HEADS, SEQ, HEAD_DIM, dtype=dtype, device='cuda:0')
torch.cuda.synchronize()
scale = HEAD_DIM ** -0.5
attn = torch.matmul(Q, K.transpose(-2, -1)) * scale
torch.cuda.synchronize()
attn_max = attn.max(dim=-1, keepdim=True).values
exp_attn = torch.exp(attn - attn_max)
attn_w = exp_attn / exp_attn.sum(dim=-1, keepdim=True)
torch.cuda.synchronize()
out = torch.matmul(attn_w, V)
torch.cuda.synchronize()
return out
test_op("Manual SDPA (24 heads, 1024 seq, 160 dim)", test_manual_sdpa)
# Test 9: Einsum (used in rope)
def test_einsum():
pos = torch.randn(B, SEQ, dtype=torch.float32, device='cuda:0')
omega = torch.randn(HEAD_DIM // 2, dtype=torch.float32, device='cuda:0')
return torch.einsum("...n,d->...nd", pos, omega)
test_op("einsum (rope-like)", test_einsum)
# Test 10: Sequential like model forward (LayerNorm → Linear → SiLU → Linear)
def test_sequential():
h = F.layer_norm(x, (DIM,), weight_ln, bias_ln)
torch.cuda.synchronize()
w1 = torch.randn(DIM*4, DIM, dtype=dtype, device='cuda:0')
h = F.linear(h, w1)
torch.cuda.synchronize()
h = F.silu(h)
torch.cuda.synchronize()
w2 = torch.randn(DIM, DIM*4, dtype=dtype, device='cuda:0')
h = F.linear(h, w2)
torch.cuda.synchronize()
return h
test_op("Sequential: LN → Linear → SiLU → Linear", test_sequential)
print("\n=== ALL TESTS PASSED ===", flush=True)