Uploaded sanitized BC250/ROCm Repository.
This commit is contained in:
@@ -0,0 +1,125 @@
|
||||
import torch
|
||||
import time
|
||||
import os
|
||||
import sys
|
||||
|
||||
print(f'PyTorch: {torch.__version__}')
|
||||
print(f'CUDA: {torch.cuda.is_available()}')
|
||||
if torch.cuda.is_available():
|
||||
print(f'Device: {torch.cuda.get_device_name(0)}')
|
||||
else:
|
||||
print('NO CUDA AVAILABLE')
|
||||
os._exit(1)
|
||||
|
||||
print()
|
||||
tests = []
|
||||
|
||||
# Test 1: Small matmul
|
||||
print('=== Test 1: Small matmul 256x256 fp32 ===')
|
||||
try:
|
||||
a = torch.randn(256, 256, device='cuda', dtype=torch.float32)
|
||||
b = torch.randn(256, 256, device='cuda', dtype=torch.float32)
|
||||
torch.cuda.synchronize()
|
||||
t = time.time()
|
||||
for _ in range(10):
|
||||
c = torch.mm(a, b)
|
||||
torch.cuda.synchronize()
|
||||
elapsed = time.time()-t
|
||||
print(f' OK: {elapsed:.3f}s')
|
||||
tests.append(('small_mm', True))
|
||||
del a, b, c
|
||||
except Exception as e:
|
||||
print(f' FAIL: {e}')
|
||||
tests.append(('small_mm', False))
|
||||
|
||||
# Test 2: Medium matmul
|
||||
print('=== Test 2: matmul 1024x1024 fp32 ===')
|
||||
try:
|
||||
a = torch.randn(1024, 1024, device='cuda', dtype=torch.float32)
|
||||
b = torch.randn(1024, 1024, device='cuda', dtype=torch.float32)
|
||||
torch.cuda.synchronize()
|
||||
t = time.time()
|
||||
c = torch.mm(a, b)
|
||||
torch.cuda.synchronize()
|
||||
elapsed = time.time()-t
|
||||
print(f' OK: {elapsed:.3f}s')
|
||||
tests.append(('med_mm', True))
|
||||
del a, b, c
|
||||
except Exception as e:
|
||||
print(f' FAIL: {e}')
|
||||
tests.append(('med_mm', False))
|
||||
|
||||
# Test 3: fp16
|
||||
print('=== Test 3: matmul 1024x1024 fp16 ===')
|
||||
try:
|
||||
a = torch.randn(1024, 1024, device='cuda', dtype=torch.float16)
|
||||
b = torch.randn(1024, 1024, device='cuda', dtype=torch.float16)
|
||||
torch.cuda.synchronize()
|
||||
t = time.time()
|
||||
c = torch.mm(a, b)
|
||||
torch.cuda.synchronize()
|
||||
elapsed = time.time()-t
|
||||
print(f' OK: {elapsed:.3f}s')
|
||||
tests.append(('fp16_mm', True))
|
||||
del a, b, c
|
||||
except Exception as e:
|
||||
print(f' FAIL: {e}')
|
||||
tests.append(('fp16_mm', False))
|
||||
|
||||
# Test 4: Conv2d
|
||||
print('=== Test 4: Conv2d 32ch fp32 ===')
|
||||
try:
|
||||
conv = torch.nn.Conv2d(32, 32, 3, padding=1).cuda().float()
|
||||
x = torch.randn(1, 32, 32, 32, device='cuda', dtype=torch.float32)
|
||||
torch.cuda.synchronize()
|
||||
t = time.time()
|
||||
y = conv(x)
|
||||
torch.cuda.synchronize()
|
||||
elapsed = time.time()-t
|
||||
print(f' OK: {elapsed:.3f}s')
|
||||
tests.append(('conv2d', True))
|
||||
del conv, x, y
|
||||
except Exception as e:
|
||||
print(f' FAIL: {e}')
|
||||
tests.append(('conv2d', False))
|
||||
|
||||
# Test 5: BMM (attention-like)
|
||||
print('=== Test 5: BMM 4x512x64 fp32 ===')
|
||||
try:
|
||||
q = torch.randn(1, 4, 512, 64, device='cuda', dtype=torch.float32)
|
||||
k = torch.randn(1, 4, 512, 64, device='cuda', dtype=torch.float32)
|
||||
torch.cuda.synchronize()
|
||||
t = time.time()
|
||||
attn = torch.matmul(q, k.transpose(-2, -1))
|
||||
torch.cuda.synchronize()
|
||||
elapsed = time.time()-t
|
||||
print(f' OK: {elapsed:.3f}s')
|
||||
tests.append(('bmm', True))
|
||||
del q, k, attn
|
||||
except Exception as e:
|
||||
print(f' FAIL: {e}')
|
||||
tests.append(('bmm', False))
|
||||
|
||||
# Test 6: Linear (typical model layer)
|
||||
print('=== Test 6: Linear 3072->3072 fp32 ===')
|
||||
try:
|
||||
lin = torch.nn.Linear(3072, 3072).cuda().float()
|
||||
x = torch.randn(1, 256, 3072, device='cuda', dtype=torch.float32)
|
||||
torch.cuda.synchronize()
|
||||
t = time.time()
|
||||
y = lin(x)
|
||||
torch.cuda.synchronize()
|
||||
elapsed = time.time()-t
|
||||
print(f' OK: {elapsed:.3f}s')
|
||||
tests.append(('linear', True))
|
||||
del lin, x, y
|
||||
except Exception as e:
|
||||
print(f' FAIL: {e}')
|
||||
tests.append(('linear', False))
|
||||
|
||||
print()
|
||||
passed = sum(1 for _, ok in tests if ok)
|
||||
print(f'=== {passed}/{len(tests)} tests passed ===')
|
||||
sys.stdout.flush()
|
||||
torch.cuda.synchronize()
|
||||
os._exit(0)
|
||||
Reference in New Issue
Block a user