diff --git a/examples/mlperf/model_train.py b/examples/mlperf/model_train.py index de6b8e9f5c..29c0f73b1e 100644 --- a/examples/mlperf/model_train.py +++ b/examples/mlperf/model_train.py @@ -246,7 +246,7 @@ def train_resnet(): if i == BENCHMARK: assert not math.isnan(loss) - median_step_time = sorted(step_times)[(BENCHMARK + 1) // 2] # in seconds + median_step_time = sorted(step_times)[BENCHMARK // 2] # in seconds estimated_total_minutes = int(median_step_time * steps_in_train_epoch * epochs / 60) print(f"Estimated training time: {estimated_total_minutes // 60}h{estimated_total_minutes % 60}m") print(f"epoch global_ops: {steps_in_train_epoch * GlobalCounters.global_ops:_}, " @@ -593,7 +593,7 @@ def train_retinanet(): if i == BENCHMARK: assert not math.isnan(loss) - median_step_time = sorted(step_times)[(BENCHMARK + 1) // 2] # in seconds + median_step_time = sorted(step_times)[BENCHMARK // 2] # in seconds estimated_total_minutes = int(median_step_time * steps_in_train_epoch * EPOCHS / 60) print(f"Estimated training time: {estimated_total_minutes // 60}h{estimated_total_minutes % 60}m") print(f"epoch global_ops: {steps_in_train_epoch * GlobalCounters.global_ops:_}, " @@ -868,7 +868,7 @@ def train_unet3d(): i += 1 if i == BENCHMARK: - median_step_time = sorted(step_times)[(BENCHMARK + 1) // 2] # in seconds + median_step_time = sorted(step_times)[BENCHMARK // 2] # in seconds estimated_total_minutes = int(median_step_time * SAMPLES_PER_EPOCH * NUM_EPOCHS / 60) print(f"Estimated training time: {estimated_total_minutes // 60}h{estimated_total_minutes % 60}m") if (TRAIN_BEAM or EVAL_BEAM) and epoch == start_epoch: break @@ -1167,7 +1167,7 @@ def train_bert(): i += 1 if i == BENCHMARK: - median_step_time = sorted(step_times)[(BENCHMARK + 1) // 2] # in seconds + median_step_time = sorted(step_times)[BENCHMARK // 2] # in seconds estimated_total_minutes = int(median_step_time * train_steps / 60) print(f"Estimated training time: {estimated_total_minutes // 60}h{estimated_total_minutes % 60}m") print(f"epoch global_ops: {train_steps * GlobalCounters.global_ops:_}, " @@ -1544,7 +1544,7 @@ def train_llama3(): mem_gb = GlobalCounters.mem_used / 1e9 gflops = GlobalCounters.global_ops / 1e9 / dev_time - mfu = ((6 * num_params * SEQLEN * GBS) / (dev_time * device_count * 2.3e15)) * 100 + mfu = ((6 * num_params * SEQLEN * GBS) / (dev_time * device_count * (4.6e15 if FP8 else 2.3e15))) * 100 tqdm.write( f"{i:5} {step_time:.3f} s step, {gbs_time:.3f} s gbs, {optim_time:.3f} s optim, {data_time:.3f} s data, {loss:.4f} loss, " \ f"{lr:.12f} LR, {grad_norm:.6f} grad_norm, {mem_gb:.2f} GB used, {gflops:9.2f} GFLOPS, {mfu:5.2f}% MFU") @@ -1577,8 +1577,9 @@ def train_llama3(): safe_save(get_state_dict(scheduler), fn) if i == BENCHMARK: - median_step_time = sorted(step_times)[(BENCHMARK + 1) // 2] - estimated_total_minutes = int(median_step_time * (SAMPLES // GBS) / 60) + median_step_time = sorted(step_times)[BENCHMARK // 2] + estimated_steps = 200_000 // GBS if getenv("LLAMA3_SIZE", "8B") == "8B" else MAX_STEPS + estimated_total_minutes = int(median_step_time * estimated_steps / 60) print(f"Estimated training time: {estimated_total_minutes // 60}h{estimated_total_minutes % 60}m") print(f"epoch global_ops: {GlobalCounters.global_ops:_}, " f"epoch global_mem: {GlobalCounters.global_mem:_}") diff --git a/examples/mlperf/models/flat_llama.py b/examples/mlperf/models/flat_llama.py index 91f2eac2da..6373888088 100644 --- a/examples/mlperf/models/flat_llama.py +++ b/examples/mlperf/models/flat_llama.py @@ -35,12 +35,15 @@ def quantize_fp8(x:Tensor, amax_state:Tensor|None=None): return x_clamped.cast(FP8_DTYPE), scale.float().reciprocal() def matmul(x:Tensor, w:Tensor, fp8=FP8, amax_x:Tensor|None=None, amax_w:Tensor|None=None) -> Tensor: - if not fp8: return x @ w.T - from tinygrad.helpers import ASM_GEMM + if not fp8: + if getenv("ASM_GEMM"): + from extra.gemm.cdna_asm_gemm import can_use_asm_gemm, asm_gemm + if can_use_asm_gemm(x, w.T): return asm_gemm(x, w.T) + return x @ w.T x_fp8, x_scale = quantize_fp8(x, amax_state=amax_x) w_fp8, w_scale = quantize_fp8(w, amax_state=amax_w) combined_scale = x_scale * w_scale - if ASM_GEMM: + if getenv("ASM_GEMM"): from extra.gemm.cdna_asm_gemm import can_use_asm_gemm, asm_gemm if can_use_asm_gemm(x_fp8, w_fp8.T): return asm_gemm(x_fp8, w_fp8.T, combined_scale=combined_scale) return x_fp8.dot(w_fp8.T, dtype=dtypes.float) * combined_scale @@ -122,7 +125,11 @@ class FlatTransformer: xq, xk = apply_rotary_emb(xq, xk, freqs_cis) if FP8: xq, xk, xv = xq.cast(dtypes.bfloat16), xk.cast(dtypes.bfloat16), xv.cast(dtypes.bfloat16) xq, xk, xv = xq.transpose(1, 2), xk.transpose(1, 2), xv.transpose(1, 2) - attn = xq.scaled_dot_product_attention(xk, xv, is_causal=True, enable_gqa=True).transpose(1, 2) + if getenv("HK_FLASH_ATTENTION"): + from extra.thunder.amd.fa import flash_attention + attn = flash_attention(xq, xk, xv, is_causal=True) + else: + attn = xq.scaled_dot_product_attention(xk, xv, is_causal=True, enable_gqa=True).transpose(1, 2) attn = attn.reshape(bsz, seqlen, -1) return matmul(attn, wo, amax_x=amax_xo, amax_w=amax_wo) @@ -193,7 +200,7 @@ class FlatTransformer: self.attention_norm[i], self.wo[i], self.ffn_norm[i], self.w1[i], self.w2[i], self.w3[i], **attn_kwargs, **amax_attn, **amax_layer) - logits = (self.norm(h).contiguous().contiguous_backward() @ self.output[0].T).contiguous_backward() + logits = matmul(self.norm(h).contiguous().contiguous_backward(), self.output[0], fp8=False).contiguous_backward() return logits def _get_pads(uop:UOp) -> list[UOp]: diff --git a/examples/mlperf/models/llama.py b/examples/mlperf/models/llama.py deleted file mode 100644 index 0ae17fd93c..0000000000 --- a/examples/mlperf/models/llama.py +++ /dev/null @@ -1,80 +0,0 @@ -from tinygrad import Tensor, nn -from tinygrad.helpers import getenv -from extra.models.llama import apply_rotary_emb, precompute_freqs_cis - -class Attention: - def __init__(self, dim:int, n_heads:int, n_kv_heads:int|None=None, linear=nn.Linear): - self.n_heads = n_heads - self.n_kv_heads = n_kv_heads if n_kv_heads is not None else n_heads # n_kv_heads != n_heads implies MQA [arxiv/2307.09288, A.2.1] - self.head_dim = dim // n_heads - self.n_rep = self.n_heads // self.n_kv_heads - - if getenv("WQKV"): - self.wqkv = linear(dim, self.n_heads * self.head_dim + self.n_kv_heads * self.head_dim * 2, bias=False) - else: - self.wq = linear(dim, self.n_heads * self.head_dim, bias=False) - self.wk = linear(dim, self.n_kv_heads * self.head_dim, bias=False) - self.wv = linear(dim, self.n_kv_heads * self.head_dim, bias=False) - - self.wo = linear(self.n_heads * self.head_dim, dim, bias=False) - - def __call__(self, x:Tensor, freqs_cis:Tensor) -> Tensor: - if getenv("WQKV"): - xqkv = self.wqkv(x) - xqkv = xqkv.reshape(xqkv.shape[0], xqkv.shape[1], self.n_kv_heads, self.n_rep + 2, self.head_dim) - xq = xqkv[:, :, :, :self.n_rep].reshape(xqkv.shape[0], xqkv.shape[1], -1) - xk = xqkv[:, :, :, self.n_rep:self.n_rep+1].reshape(xqkv.shape[0], xqkv.shape[1], -1) - xv = xqkv[:, :, :, self.n_rep+1:self.n_rep+2].reshape(xqkv.shape[0], xqkv.shape[1], -1) - else: - xq, xk, xv = self.wq(x), self.wk(x), self.wv(x) - - xq = xq.reshape(xq.shape[0], xq.shape[1], self.n_heads, self.head_dim) - xk = xk.reshape(xk.shape[0], xk.shape[1], self.n_kv_heads, self.head_dim) - xv = xv.reshape(xv.shape[0], xv.shape[1], self.n_kv_heads, self.head_dim) - - xq, xk = apply_rotary_emb(xq, xk, freqs_cis) - bsz, seqlen, _, _ = xq.shape - - xq, xk, xv = xq.transpose(1, 2), xk.transpose(1, 2), xv.transpose(1, 2) - attn = xq.scaled_dot_product_attention(xk, xv, is_causal=True, enable_gqa=True).transpose(1, 2) - - attn = attn.reshape(bsz, seqlen, -1) - return self.wo(attn) - -class FeedForward: - def __init__(self, dim:int, hidden_dim:int, linear=nn.Linear): - self.w1 = linear(dim, hidden_dim, bias=False) - self.w2 = linear(hidden_dim, dim, bias=False) - self.w3 = linear(dim, hidden_dim, bias=False) # the gate in Gated Linear Unit - - def __call__(self, x:Tensor) -> Tensor: - w1 = self.w1(x).silu() - w3 = self.w3(x) - return self.w2(w1 * w3) - -class TransformerBlock: - def __init__(self, dim:int, hidden_dim:int, n_heads:int, n_kv_heads:int|None, norm_eps:float, linear=nn.Linear): - self.attention = Attention(dim, n_heads, n_kv_heads, linear) - self.feed_forward = FeedForward(dim, hidden_dim, linear) - self.attention_norm = nn.RMSNorm(dim, norm_eps) - self.ffn_norm = nn.RMSNorm(dim, norm_eps) - - def __call__(self, x:Tensor, freqs_cis:Tensor): - h = x + self.attention(self.attention_norm(x), freqs_cis) - return h + self.feed_forward(self.ffn_norm(h)) - -class Transformer: - def __init__(self, dim:int, hidden_dim:int, n_heads:int, n_layers:int, norm_eps:float, vocab_size:int, n_kv_heads:int|None=None, - rope_theta:int=10000, max_context:int=1024, linear=nn.Linear, embedding=nn.Embedding): - self.layers = [TransformerBlock(dim, hidden_dim, n_heads, n_kv_heads, norm_eps, linear) for _ in range(n_layers)] - self.norm = nn.RMSNorm(dim, norm_eps) - self.tok_embeddings = embedding(vocab_size, dim) - self.output = nn.Linear(dim, vocab_size, bias=False) if embedding == nn.Embedding else linear(dim, vocab_size, bias=False) - self.freqs_cis = precompute_freqs_cis(dim // n_heads, max_context * 2, rope_theta).contiguous().requires_grad_(False) - - def __call__(self, tokens:Tensor): - h = self.tok_embeddings(tokens) - freqs_cis = self.freqs_cis.cast(h.dtype)[:, :tokens.shape[1], :, :, :] - for layer in self.layers: h = layer(h, freqs_cis) - logits = self.output(self.norm(h)) - return logits diff --git a/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/dev_beam.sh b/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/dev_beam.sh index c345538f6f..6dfbf57089 100755 --- a/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/dev_beam.sh +++ b/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/dev_beam.sh @@ -36,7 +36,7 @@ export DATA_SEED=${DATA_SEED:-5760} export JITBEAM=${JITBEAM:-3} export BEAM_UOPS_MAX=6000 BEAM_UPCAST_MAX=256 BEAM_LOCAL_MAX=1024 BEAM_MIN_PROGRESS=5 BEAM_PADTO=1 -export FAKEDATA=1 BENCHMARK=10 +export FAKEDATA=1 BENCHMARK=${BENCHMARK:-10} if [ -z "$FULL_LAYERS" ]; then export LLAMA_LAYERS=2 fi diff --git a/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/profile.sh b/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/profile.sh index cfddfa2601..e8fed36b18 100755 --- a/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/profile.sh +++ b/examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/profile.sh @@ -1,5 +1,5 @@ #!/bin/bash export BENCHMARK=5 export EVAL_BS=0 -VIZ=${VIZ:--1} examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/dev_run.sh +VIZ=${VIZ:--1} FULL_LAYERS=1 DEBUG=0 examples/mlperf/training_submission_v6.0/tinycorp/benchmarks/llama8b/implementations/tinybox_8xMI350X/dev_beam.sh extra/viz/cli.py --profile -s "${DEV:-AMD}" diff --git a/extra/gemm/cdna_asm_gemm.py b/extra/gemm/cdna_asm_gemm.py index 42524183d4..cbeb8995df 100644 --- a/extra/gemm/cdna_asm_gemm.py +++ b/extra/gemm/cdna_asm_gemm.py @@ -2716,8 +2716,11 @@ def custom_gemm_bw(gradient:UOp, kernel:UOp): assert all_same([gradient.device, a.device, b.device, out.device]) a_t, b_t, g_t = Tensor(a, device=a.device), Tensor(b, device=a.device), Tensor(gradient, device=a.device) g_t = g_t[:a.shape[0]] - grad_a = (g_t @ b_t.T).uop - grad_b = (a_t.permute(2, 0, 1).reshape(a_t.shape[2], -1) @ g_t.reshape(-1, g_t.shape[-1])).uop + if can_use_asm_gemm(g_t, b_t.T): grad_a = asm_gemm(g_t, b_t.T).uop + else: grad_a = (g_t @ b_t.T).uop + a_t_flat, g_t_flat = a_t.permute(2, 0, 1).reshape(a_t.shape[2], -1), g_t.reshape(-1, g_t.shape[-1]) + if can_use_asm_gemm(a_t_flat, g_t_flat): grad_b = asm_gemm(a_t_flat, g_t_flat).uop + else: grad_b = (a_t_flat @ g_t_flat).uop return (None, grad_a, grad_b) # ** main gemm function diff --git a/extra/gemm/rdna4_asm_matmul.py b/extra/gemm/rdna4_asm_matmul.py new file mode 100644 index 0000000000..04ed8ce908 --- /dev/null +++ b/extra/gemm/rdna4_asm_matmul.py @@ -0,0 +1,245 @@ +# RDNA4 128x128 GEMM using WMMA — optimized DS scheduling +import numpy as np +from tinygrad import Tensor, Device, Context, GlobalCounters +from tinygrad.uop.ops import UOp, Ops, KernelInfo +from tinygrad.helpers import getenv, colored +from tinygrad.dtype import dtypes, AddrSpace +from tinygrad.engine.realize import Estimates +from tinygrad.renderer.amd.dsl import s, v, VCC_LO, NULL, src, ttmp +from tinygrad.runtime.autogen.amd.rdna4.ins import * + +BLOCK_M, BLOCK_N, BLOCK_K = 128, 128, 16 +TILES_M, TILES_N = 4, 4 +THREADS, ELEM = 128, 2 +LDS_A_ROW = BLOCK_K*ELEM # 32 +LDS_B_ROW = BLOCK_N*ELEM # 256 +LDS_A_SIZE = BLOCK_M * LDS_A_ROW # 4096 +LDS_B_SIZE = BLOCK_K * LDS_B_ROW # 4096 +LDS_SIZE = LDS_A_SIZE + LDS_B_SIZE # 8192 +LDS_B_OFF = LDS_A_SIZE +ACC, DA, DB, FA, FB, ET = 60, 188, 196, 204, 44, 10 + +def build_kernel(N, arch='gfx1200'): + assert N % BLOCK_M == 0 and N >= 256 + NO_ALU, NO_DS, NO_GLOBAL = getenv("NO_ALU", 0), getenv("NO_DS", 0), getenv("NO_GLOBAL", 0) + I, L, B = [], {}, [] + def e(i): I.append(i); return i + def label(n): L[n] = sum(i.size() for i in I) + def br(i, t): B.append((len(I)-1, t)) + + e(s_load_b128(sdata=s[4:7], sbase=s[0:1], ioffset=0, soffset=NULL)) + e(s_load_b64(sdata=s[8:9], sbase=s[0:1], ioffset=0x10, soffset=NULL)) + e(s_wait_kmcnt(simm16=0)) + e(s_mov_b32(s[10], ttmp[9])); e(s_and_b32(s[11], ttmp[7], 0xFFFF)) + e(s_lshl_b32(s[10], s[10], 7)); e(s_lshl_b32(s[11], s[11], 7)) + e(s_mov_b32(s[12], N)); e(s_lshl_b32(s[13], s[12], 1)) + e(s_mul_i32(s[14], s[12], BLOCK_K*ELEM)) + e(s_add_co_i32(s[17], s[12], -2*BLOCK_K)) # loop bound + + e(v_and_b32_e32(v[1], 31, v[0])); e(v_lshrrev_b32_e32(v[2], 5, v[0])) + e(v_and_b32_e32(v[3], 1, v[2])); e(v_lshrrev_b32_e32(v[2], 1, v[2])) + + e(v_lshlrev_b32_e32(v[4], 5, v[0])) + # B store: transposed layout for stride-32 reads. addr = LDS_B_OFF + (tid%8)*512 + (tid/8)*32 + e(v_and_b32_e32(v[48], 7, v[0])); e(v_lshlrev_b32_e32(v[5], 9, v[48])) # (tid%8)*512 + e(v_lshrrev_b32_e32(v[48], 3, v[0])); e(v_lshlrev_b32_e32(v[48], 5, v[48])) # (tid/8)*32 + e(v_add_nc_u32_e32(v[5], v[5], v[48])); e(v_add_nc_u32_e32(v[5], LDS_B_OFF, v[5])) + + e(v_add_nc_u32_e32(v[48], s[11], v[0])) + e(v_mul_lo_u32(v[6], v[48], N*ELEM)); e(v_mov_b32_e32(v[7], 0)) + e(v_lshrrev_b32_e32(v[48], 3, v[0])); e(v_mul_lo_u32(v[8], v[48], N*ELEM)) + e(v_and_b32_e32(v[48], 7, v[0])); e(v_lshlrev_b32_e32(v[48], 5, v[48])) + e(v_add_nc_u32_e32(v[8], v[8], v[48])) + e(s_mul_i32(s[15], s[10], ELEM)); e(v_add_nc_u32_e32(v[8], s[15], v[8])) + e(v_mov_b32_e32(v[9], 0)) + + # LDS read addrs with padded strides (eliminates bank conflicts) + # A: (lane%16)*LDS_A_ROW + (lane/16)*16 + wave_m*64*LDS_A_ROW + # B: (lane%16)*LDS_B_ROW + (lane/16)*16 + wave_n*64*ELEM + LDS_B_OFF + LLA, LLB = 40, 43 + e(v_and_b32_e32(v[50], 15, v[1])); e(v_lshrrev_b32_e32(v[51], 4, v[1])) + e(v_lshlrev_b32_e32(v[LLA], 5, v[50])) # (lane%16) * 32 + e(v_lshlrev_b32_e32(v[51], 4, v[51])) # (lane/16) * 16 + e(v_add_nc_u32_e32(v[LLA], v[LLA], v[51])) + e(v_lshlrev_b32_e32(v[52], 11, v[2])) # wave_m * 2048 + e(v_add_nc_u32_e32(v[LLA], v[LLA], v[52])) + # B read: transposed layout. addr = LDS_B_OFF + (lane%16)*32 + (lane/16)*16 + wave_n*2*512 + # wave_n selects column panels: wave_n*2 panels (each panel=16 cols, wave_n covers 64 cols = 4 panels) + # But wave_n*2*512 = wave_n*1024. Hmm, wave_n covers cols [wave_n*64 : (wave_n+1)*64]. + # Each panel = 16 cols = 512 bytes. wave_n*64/16 = wave_n*4 panels. Offset = wave_n*4*512 = wave_n*2048. + e(v_lshlrev_b32_e32(v[LLB], 5, v[50])) # (lane%16) * 32 (stride 32!) + e(v_add_nc_u32_e32(v[LLB], v[LLB], v[51])) # + (lane/16)*16 + e(v_lshlrev_b32_e32(v[52], 11, v[3])) # wave_n * 2048 + e(v_add_nc_u32_e32(v[LLB], v[LLB], v[52])) + e(v_add_nc_u32_e32(v[LLB], LDS_B_OFF, v[LLB])) + + for i in range(0, 128, 2): + e(VOPD(VOPDOp.V_DUAL_MOV_B32, VOPDOp.V_DUAL_MOV_B32, vdstx=v[ACC+i], vdsty=v[ACC+i+1], srcx0=0, srcy0=0)) + e(s_mov_b32(s[16], 0)) + + if not NO_GLOBAL: + for i in range(2): e(global_load_b128(vdst=v[DA+i*4:DA+i*4+3], vaddr=v[6:7], saddr=s[4:5], ioffset=i*16)) + for i in range(2): e(global_load_b128(vdst=v[DB+i*4:DB+i*4+3], vaddr=v[8:9], saddr=s[6:7], ioffset=i*16)) + e(s_wait_loadcnt(simm16=0)) + if not NO_DS: + for i in range(2): e(ds_store_b128(addr=v[4], data0=v[DA+i*4:DA+i*4+3], offset0=(i*16)&0xFF, offset1=(i*16)>>8)) + for i in range(2): e(ds_store_b128(addr=v[5], data0=v[DB+i*4:DB+i*4+3], offset0=(i*16)&0xFF, offset1=(i*16)>>8)) + if not NO_GLOBAL: + e(v_add_nc_u32_e32(v[6], BLOCK_K*ELEM, v[6])) + e(v_add_nc_u32_e32(v[8], s[14], v[8])) + + # ============================================================================= + def emit_iter_body(load_set='AB'): + if not NO_DS: + e(s_wait_dscnt(simm16=0)) + e(s_barrier_signal(ssrc0=src[193])); e(s_barrier_wait(simm16=0xFFFF)) + if not NO_GLOBAL: + if 'A' in load_set: + for i in range(2): e(global_load_b128(vdst=v[DA+i*4:DA+i*4+3], vaddr=v[6:7], saddr=s[4:5], ioffset=i*16)) + e(v_add_nc_u32_e32(v[6], BLOCK_K*ELEM, v[6])) + if 'B' in load_set: + for i in range(2): e(global_load_b128(vdst=v[DB+i*4:DB+i*4+3], vaddr=v[8:9], saddr=s[6:7], ioffset=i*16)) + e(v_add_nc_u32_e32(v[8], s[14], v[8])) + if not NO_DS: + # Issue 6 loads: A[0:3] + B[0] + B[1]. B[2:3] interleaved with WMMAs. + for tm in range(TILES_M): + aoff = tm * 16 * LDS_A_ROW + e(ds_load_b128(vdst=v[FA+tm*4:FA+tm*4+3], addr=v[LLA], offset0=aoff&0xFF, offset1=aoff>>8)) + e(ds_load_b128(vdst=v[FB:FB+3], addr=v[LLB], offset0=0, offset1=0)) + e(ds_load_b128(vdst=v[FB+4:FB+7], addr=v[LLB], offset0=0, offset1=2)) + e(s_wait_dscnt(simm16=0)) # wait for 6 loads (no stall!) + if not NO_ALU: + # B[0] WMMAs — issue B[2] during compute + if not NO_DS: e(ds_load_b128(vdst=v[FB+8:FB+11], addr=v[LLB], offset0=0, offset1=4)) + for tm in range(TILES_M): + ac = ACC + (tm*TILES_N+0)*8 + e(v_wmma_f32_16x16x16_f16(vdst=v[ac:ac+7], src0=v[FA+tm*4:FA+tm*4+3], src1=v[FB:FB+3], src2=v[ac:ac+7])) + # B[1] WMMAs — issue B[3] during compute + if not NO_DS: + e(ds_load_b128(vdst=v[FB+12:FB+15], addr=v[LLB], offset0=0, offset1=6)) + for tm in range(TILES_M): + ac = ACC + (tm*TILES_N+1)*8 + e(v_wmma_f32_16x16x16_f16(vdst=v[ac:ac+7], src0=v[FA+tm*4:FA+tm*4+3], src1=v[FB+4:FB+7], src2=v[ac:ac+7])) + # B[2] WMMAs — B[2] loaded during B[0] WMMAs (~100 cycles ago) + if not NO_DS: e(s_wait_dscnt(simm16=1)) # B[2] done, B[3] may still be loading + for tm in range(TILES_M): + ac = ACC + (tm*TILES_N+2)*8 + e(v_wmma_f32_16x16x16_f16(vdst=v[ac:ac+7], src0=v[FA+tm*4:FA+tm*4+3], src1=v[FB+8:FB+11], src2=v[ac:ac+7])) + # B[3] WMMAs + if not NO_DS: e(s_wait_dscnt(simm16=0)) + for tm in range(TILES_M): + ac = ACC + (tm*TILES_N+3)*8 + e(v_wmma_f32_16x16x16_f16(vdst=v[ac:ac+7], src0=v[FA+tm*4:FA+tm*4+3], src1=v[FB+12:FB+15], src2=v[ac:ac+7])) + if not NO_GLOBAL and not NO_DS: e(s_wait_loadcnt(simm16=0)) + if not NO_DS: + for i in range(2): e(ds_store_b128(addr=v[4], data0=v[DA+i*4:DA+i*4+3], offset0=(i*16)&0xFF, offset1=(i*16)>>8)) + for i in range(2): e(ds_store_b128(addr=v[5], data0=v[DB+i*4:DB+i*4+3], offset0=(i*16)&0xFF, offset1=(i*16)>>8)) + e(s_add_co_i32(s[16], s[16], BLOCK_K)) + + label('LOOP') + emit_iter_body(load_set='A') + emit_iter_body(load_set='B') + e(s_cmp_lt_i32(s[16], s[17])); e(s_cbranch_scc1(simm16=0)); br(I[-1], 'LOOP') + + emit_iter_body(load_set='AB') # tail with prefetch + + # Final iteration: no prefetch, no ds_store needed + if not NO_DS: + e(s_wait_dscnt(simm16=0)) + e(s_barrier_signal(ssrc0=src[193])); e(s_barrier_wait(simm16=0xFFFF)) + if not NO_DS: + for tm in range(TILES_M): + aoff = tm * 16 * LDS_A_ROW + e(ds_load_b128(vdst=v[FA+tm*4:FA+tm*4+3], addr=v[LLA], offset0=aoff&0xFF, offset1=aoff>>8)) + e(ds_load_b128(vdst=v[FB:FB+3], addr=v[LLB], offset0=0, offset1=0)) + e(ds_load_b128(vdst=v[FB+4:FB+7], addr=v[LLB], offset0=0, offset1=2)) + e(s_wait_dscnt(simm16=0)) + if not NO_ALU: + if not NO_DS: e(ds_load_b128(vdst=v[FB+8:FB+11], addr=v[LLB], offset0=0, offset1=4)) + for tm in range(TILES_M): + ac = ACC + (tm*TILES_N+0)*8 + e(v_wmma_f32_16x16x16_f16(vdst=v[ac:ac+7], src0=v[FA+tm*4:FA+tm*4+3], src1=v[FB:FB+3], src2=v[ac:ac+7])) + if not NO_DS: e(ds_load_b128(vdst=v[FB+12:FB+15], addr=v[LLB], offset0=0, offset1=6)) + for tm in range(TILES_M): + ac = ACC + (tm*TILES_N+1)*8 + e(v_wmma_f32_16x16x16_f16(vdst=v[ac:ac+7], src0=v[FA+tm*4:FA+tm*4+3], src1=v[FB+4:FB+7], src2=v[ac:ac+7])) + if not NO_DS: e(s_wait_dscnt(simm16=1)) + for tm in range(TILES_M): + ac = ACC + (tm*TILES_N+2)*8 + e(v_wmma_f32_16x16x16_f16(vdst=v[ac:ac+7], src0=v[FA+tm*4:FA+tm*4+3], src1=v[FB+8:FB+11], src2=v[ac:ac+7])) + if not NO_DS: e(s_wait_dscnt(simm16=0)) + for tm in range(TILES_M): + ac = ACC + (tm*TILES_N+3)*8 + e(v_wmma_f32_16x16x16_f16(vdst=v[ac:ac+7], src0=v[FA+tm*4:FA+tm*4+3], src1=v[FB+12:FB+15], src2=v[ac:ac+7])) + + label('EPILOGUE') + e(v_and_b32_e32(v[ET], 15, v[1])) + e(v_lshrrev_b32_e32(v[ET+1], 4, v[1])); e(v_lshlrev_b32_e32(v[ET+1], 3, v[ET+1])) + e(v_lshlrev_b32_e32(v[ET+2], 6, v[2])); e(v_add_nc_u32_e32(v[ET+2], s[11], v[ET+2])) + e(v_lshlrev_b32_e32(v[ET+3], 6, v[3])); e(v_add_nc_u32_e32(v[ET+3], s[10], v[ET+3])) + e(v_add_nc_u32_e32(v[ET+3], v[ET+3], v[ET])); e(v_mov_b32_e32(v[ET+5], 0)) + + for tm in range(TILES_M): + for tn in range(TILES_N): + ac = ACC + (tm*TILES_N+tn)*8; r_off, c_off = tm*16, tn*16 + e(v_add_nc_u32_e32(v[ET+6], r_off, v[ET+2])); e(v_add_nc_u32_e32(v[ET+6], v[ET+1], v[ET+6])) + e(v_mul_lo_u32(v[ET+4], v[ET+6], s[12])); e(v_add_nc_u32_e32(v[ET+4], v[ET+4], v[ET+3])) + if c_off: e(v_add_nc_u32_e32(v[ET+4], c_off, v[ET+4])) + e(v_lshlrev_b32_e32(v[ET+4], 1, v[ET+4])) + for elem in range(8): + e(v_cvt_f16_f32_e32(v[ET+7], v[ac+elem])) + e(global_store_b16(vaddr=v[ET+4:ET+5], vsrc=v[ET+7], saddr=s[8:9])) + if elem < 7: e(v_add_nc_u32_e32(v[ET+4], s[13], v[ET+4])) + + e(s_wait_storecnt(simm16=0)); e(s_sendmsg(simm16=3)); e(s_endpgm()) + + for idx, target in B: + off = (L[target] - sum(i.size() for i in I[:idx+1])) // 4 + assert -32768 <= off <= 32767; I[idx].simm16 = off + return I + +N = getenv("N", 4096) + +def test_matmul(): + dev = Device[Device.DEFAULT] + arch = getattr(dev.renderer, 'arch', 'gfx1200') + print(f"Device arch: {arch}") + insts = build_kernel(N, arch) + + rng = np.random.default_rng(42) + a = Tensor(rng.random((N, N), dtype=np.float32).astype(np.float16)) + b = Tensor(rng.random((N, N), dtype=np.float32).astype(np.float16)) + c = Tensor.empty(N, N, dtype=dtypes.half) + Tensor.realize(a, b, c) + + grid, local = (N//BLOCK_N, N//BLOCK_M, 1), (THREADS, 1, 1) + print(f"Grid: {grid}, Local: {local}") + + dname = Device.DEFAULT + def asm_kernel(A, B, C): + gidxs = [UOp.special(n, f"gidx{i}") for i,n in enumerate(grid)] + lidxs = [UOp.special(THREADS, "lidx0")] + lds = UOp(Ops.DEFINE_LOCAL, dtypes.uint8.ptr(size=max(LDS_SIZE, 65536//getenv("LIMIT_OCC",2)), addrspace=AddrSpace.LOCAL), (), 'lds') + sink = UOp.sink(A.base, B.base, C.base, lds, *gidxs, *lidxs, + arg=KernelInfo(name=colored("kernel","cyan"), estimates=Estimates(ops=N*N*N*2, mem=N*N*2*3))) + return UOp(Ops.PROGRAM, src=(sink, UOp(Ops.DEVICE, arg=dname), UOp(Ops.LINEAR, src=tuple([UOp(Ops.INS, arg=x) for x in insts])))) + + c = Tensor.custom_kernel(a, b, c, fxn=asm_kernel)[2] + ei = c.schedule()[0].lower() + + ets = [] + with Context(DEBUG=2): + for _ in range(getenv("CNT", 5)): ets.append(ei.run(wait=True)) + print(f"REAL TFLOPS {N*N*N*2 / min(ets) * 1e-12:.2f}") + + if getenv("VERIFY", 1): + GlobalCounters.reset() + c_np = c.float().numpy() + a_np, b_np = a.float().numpy(), b.float().numpy() + ref = a_np @ b_np + err = np.sqrt(np.mean((c_np - ref)**2)) / np.sqrt(np.mean(ref**2)) + print(f"relative RMSE {err:.6f}") + if err != err or err > 0.05: raise RuntimeError(f"matmul is wrong! RMSE={err}") + +if __name__ == "__main__": + test_matmul() diff --git a/extra/mlx_driver/loopback.py b/extra/mlx_driver/loopback.py index 2013263f9f..31a017f061 100644 --- a/extra/mlx_driver/loopback.py +++ b/extra/mlx_driver/loopback.py @@ -6,7 +6,8 @@ from tinygrad.device import Device, BufferSpec from tinygrad.runtime.support.system import PCIDevice from tinygrad.runtime.support.memory import AddrSpace from tinygrad.runtime.ops_amd import AMDComputeQueue -from extra.mlx_driver.mlxdev import MLXDev, MLXQP, to_be +from tinygrad.helpers import to_be32, to_be64 +from extra.mlx_driver.mlxdev import MLXDev, MLXQP BUF_SIZE = 0x1000 MLX_PCI = getenv("MLX_PCI", "0000:41:00.0") @@ -49,7 +50,7 @@ rq_wqe = qp.qp_buf.view((qp.rq_head & rq_mask) * 16, 16) rq_wqe[:] = struct.pack('>IIQ', len(test_msg), dev.mkey, dst_paddr) qp.rq_head += 1 # ring recv doorbell from CPU (DBR offset 0 = recv counter) -dev.dbr[qp.qp_dbr // 4] = to_be('I', qp.rq_head) +dev.dbr[qp.qp_dbr // 4] = to_be32(qp.rq_head) # build send WQE in SQ from CPU (opcode 0x0a = SEND, ds_count=2) sq_head = qp.sq_head @@ -60,7 +61,7 @@ wqe[0:8] = struct.pack('>II', (sq_head << 8) | 0x0a, (qp.qp_info['qpn'] << 8) | wqe[11] = 0x08 # CE: signal completion wqe[16:32] = struct.pack('>IIQ', len(test_msg), dev.mkey, src_paddr) qp.sq_head += 1 -doorbell_val = to_be('Q', int.from_bytes(bytes(wqe[0:8]), 'big')) +doorbell_val = to_be64(int.from_bytes(bytes(wqe[0:8]), 'big')) # map MLX5 UAR and DBR into GPU VA uar_paddr = dev.pci_dev.bar_info(0)[0] + dev.uar * 0x1000 @@ -72,7 +73,7 @@ print(f"UAR gpu_va=0x{uar_gpu_va:x} DBR gpu_va=0x{dbr_gpu_va:x}") q = AMDComputeQueue(gpu) q.wait(gpu.timeline_signal, gpu.timeline_value - 1) # write DBR (32-bit sq_head) - send doorbell at qp_dbr + 4 -q.release_mem(dbr_gpu_va + qp.qp_dbr + 4, to_be('I', qp.sq_head), q.pm4.data_sel__mec_release_mem__send_32_bit_low, +q.release_mem(dbr_gpu_va + qp.qp_dbr + 4, to_be32(qp.sq_head), q.pm4.data_sel__mec_release_mem__send_32_bit_low, q.pm4.int_sel__mec_release_mem__none) # write UAR doorbell (64-bit) q.release_mem(uar_gpu_va + 0x800, doorbell_val, q.pm4.data_sel__mec_release_mem__send_64_bit_data, diff --git a/extra/sqtt/examples/generate_examples.py b/extra/sqtt/examples/generate_examples.py index 63eab22d7f..55ea70e4c4 100644 --- a/extra/sqtt/examples/generate_examples.py +++ b/extra/sqtt/examples/generate_examples.py @@ -10,6 +10,7 @@ EXAMPLES = { "plus":"test/test_tiny.py TestTiny.test_plus", "gemm":"-c \"from tinygrad import Tensor; (Tensor.empty(N:=32, N)@Tensor.empty(N, N)).realize()\"", "sync":"test/amd/test_custom_kernel.py TestCustomKernel.test_lds_sync", + "handwritten":"test/amd/test_custom_kernel.py TestCustomKernel.test_handwritten", } if __name__ == "__main__": @@ -21,6 +22,6 @@ if __name__ == "__main__": for i in range(2): # AM_RESET=1 gets a clear trace, does not work on mi300 machines subprocess.run([sys.executable, *shlex.split(test)], cwd=EXAMPLES_DIR.parent.parent.parent, - env={**os.environ, "AMD":"1", "AM_RESET":"1" if not arch.startswith("gfx9") else "0", "VIZ":"-2", "PYTHONPATH":"."}) + env={**os.environ, "DEV":"AMD", "AM_RESET":"1" if not arch.startswith("gfx9") else "0", "VIZ":"-2", "PYTHONPATH":"."}) PROFILE_PATH.rename(dest:=EXAMPLES_DIR/arch/f"profile_{name}_run_{i}.pkl") print(f"saved SQTT trace to {dest}") diff --git a/extra/sqtt/examples/gfx1200/profile_handwritten_run_0.pkl b/extra/sqtt/examples/gfx1200/profile_handwritten_run_0.pkl new file mode 100644 index 0000000000..f0c862b6b9 Binary files /dev/null and b/extra/sqtt/examples/gfx1200/profile_handwritten_run_0.pkl differ diff --git a/extra/sqtt/examples/gfx1200/profile_handwritten_run_1.pkl b/extra/sqtt/examples/gfx1200/profile_handwritten_run_1.pkl new file mode 100644 index 0000000000..8f8dc9f114 Binary files /dev/null and b/extra/sqtt/examples/gfx1200/profile_handwritten_run_1.pkl differ diff --git a/extra/viz/cli.py b/extra/viz/cli.py index 2fa1459372..4077976ffd 100755 --- a/extra/viz/cli.py +++ b/extra/viz/cli.py @@ -59,7 +59,6 @@ def main(args) -> None: events:list = viz.load_pickle(args.profile_path, default=[]) if (profile_bytes:=viz.get_profile(events)) is None: raise RuntimeError(f"empty profile in {args.profile_path}") profile = decode_profile(profile_bytes) - viz.load_amd_counters(viz.ctxs, events) profile["layout"].update([(f'{c["name"]} {s["name"]}', s["data"]) for c in viz.ctxs if c["name"].startswith("SQTT") for s in c["steps"] if "PKTS" in s["name"]]) if args.src is None: @@ -88,7 +87,7 @@ def main(args) -> None: op_str = hex_colored(op_name, color) if color and not args.no_color else op_name phase, delay = None, 0 idx = next(pkt_idxs.setdefault(e.device, itertools.count())) - if e.device.startswith("WAVE") or e.device == "OTHER_SIMD": + if e.device.startswith("WAVE"): inst = f"0x{(pc:=int(info.replace('PC:', ''))):05x} {pc_map[pc]}" if info else f"{'':7} {op_name}" dispatch_to_inst[f"{e.device}-{idx}"] = (inst, int(e.st)) phase = "DISPATCH" diff --git a/test/amd/test_custom_kernel.py b/test/amd/test_custom_kernel.py index 47c37b8068..ac4dad03d8 100644 --- a/test/amd/test_custom_kernel.py +++ b/test/amd/test_custom_kernel.py @@ -104,20 +104,21 @@ def custom_handwritten(A:UOp, arch:str) -> UOp: lds = UOp(Ops.DEFINE_LOCAL, dtypes.uint8.ptr(size=512, addrspace=AddrSpace.LOCAL), (), 'lds') # 128 * 4 bytes k = Kernel(arch) k.emit(r4.s_nop(0)) - k.emit(r4.v_mov_b32_e32(v[1], 10)) + k.emit(r4.v_mov_b32_e32(v[1], 4)) def emit_alt(): - for i in range(4): + for i in range(2): k.emit(r4.v_mov_b32_e32(v[20+i], 4.0)) k.emit(r4.v_rcp_f32_e32(v[22+i], v[20+i])) k.emit(r4.s_mov_b32(s[20+i], i)) + k.emit(r4.s_mul_i32(s[14+i], s[12+i], 32)) def emit_wmma(): - for _ in range(4): + for _ in range(2): k.emit(r4.v_wmma_f32_16x16x16_f16(v[0:7], v[8:11], v[8:11], 1)) k.label("start") k.emit(s_mov_b32(s[1], 10)) k.label("loop") # wmma should've overlapped here if it was a different unit? - for _ in range(4): + for _ in range(2): emit_wmma() emit_alt() for _ in range(8): k.emit(s_nop(1)) diff --git a/test/amd/test_sqttmap.py b/test/amd/test_sqttmap.py index 99ea7bcb5a..d6225e61fa 100644 --- a/test/amd/test_sqttmap.py +++ b/test/amd/test_sqttmap.py @@ -100,9 +100,7 @@ class TestSQTTMapBase(unittest.TestCase): elif "WAVE" in e.device: # sopk/immediates don't get ALU/MEM EXEC if e.name.display_name not in {"IMMEDIATE", "IMMEDIATE_MASK", "JUMP", "JUMP_NO", "MESSAGE", "BARRIER", "BARRIER_SIGNAL", - "WAVEEND", "WAVERDY"}: insts += 1 - # OTHER_ is its own stream, it's the INST from other SIMDs that share the same EXEC. - elif e.device.startswith("OTHER"): continue + "WAVEEND", "WAVERDY"} and not e.name.display_name.startswith("OTHER_"): insts += 1 else: raise Exception(f"timeline row must be INST or EXEC, got {e.device}") self.assertEqual(execs, insts) @@ -131,7 +129,18 @@ class TestSQTTMapBase(unittest.TestCase): class TestSQTTMapRDNA3(TestSQTTMapBase): target = "gfx1100" -class TestSQTTMapRDNA4(TestSQTTMapBase): target = "gfx1200" +class TestSQTTMapRDNA4(TestSQTTMapBase): + target = "gfx1200" + + @unittest.expectedFailure + def test_rdna4_wmma(self): + events, kernels, target = self.examples["profile_handwritten_run_0"] + row_ends = {} + for e in sqtt_timeline(events[0].blob, list(kernels.values())[0].lib, target): + if type(e).__name__ != "ProfileRangeEvent" or e.device != "ALUEXEC:0 WMMA": continue + if (et:=row_ends.get(e.device)) is not None and e.st < et: + raise RuntimeError(f"WMMA exec overlaps in {e.device}: {e.st} {et}.") + row_ends[e.device] = e.en class TestSQTTMapCDNA(TestSQTTMapBase): target = "gfx950" diff --git a/test/backend/test_asm_gemm.py b/test/backend/test_asm_gemm.py index b3f825caf2..ab91dc37e0 100644 --- a/test/backend/test_asm_gemm.py +++ b/test/backend/test_asm_gemm.py @@ -21,9 +21,8 @@ def run_asm_gemm(a_shape, b_shape, dtype=dtypes.float16, a_shard=None, b_shard=N a, b = a_rand.clone().requires_grad_(), b_rand.clone().requires_grad_() if multi: a, b = a.shard(devs, axis=a_shard), b.shard(devs, axis=b_shard) - with Context(ASM_GEMM=1): - tst = asm_gemm(a, b) - tst.sum().backward() + tst = asm_gemm(a, b) + tst.sum().backward() Tensor.realize(tst, a.grad, b.grad) a_ref, b_ref = a_rand.clone().requires_grad_(), b_rand.clone().requires_grad_() @@ -32,9 +31,8 @@ def run_asm_gemm(a_shape, b_shape, dtype=dtypes.float16, a_shard=None, b_shard=N a_ref = a_ref.cast(dtypes.bfloat16) b_ref = b_ref.cast(dtypes.bfloat16) if multi: a_ref, b_ref = a_ref.shard(devs, axis=a_shard), b_ref.shard(devs, axis=b_shard) - with Context(ASM_GEMM=0): - ref = asm_gemm(a_ref, b_ref) - ref.sum().backward() + ref = a_ref @ b_ref + ref.sum().backward() Tensor.realize(ref, a_ref.grad, b_ref.grad) # no validation on the NULL device @@ -136,14 +134,12 @@ class TestGemmLlama(unittest.TestCase): if not is_cdna4() or getenv("MOCKGPU"): self.skipTest("very slow on non mi350x") - @Context(ASM_GEMM=1) - def test_empty(self): (Tensor.empty(N:=getenv("N", 4096), N, dtype=self.dtype)@Tensor.empty(N, N, dtype=self.dtype)).realize() + def test_empty(self): asm_gemm(Tensor.empty(N:=getenv("N", 4096), N, dtype=self.dtype), Tensor.empty(N, N, dtype=self.dtype)).realize() - @Context(ASM_GEMM=1) def test_empty_bw(self): x = Tensor.empty(1, N:=getenv("N", 4096), N, dtype=self.dtype, requires_grad=True) y = Tensor.empty((N, N), dtype=self.dtype, requires_grad=True) - z = x @ y + z = asm_gemm(x, y) z.sum().backward() Tensor.realize(z, x.grad, y.grad) # FP8 forward output is bf16, gradients use fp8e5m2 (aka bf8) diff --git a/test/external/external_llm_eval.py b/test/external/external_llm_eval.py index 3a111b71f3..0841b9740e 100644 --- a/test/external/external_llm_eval.py +++ b/test/external/external_llm_eval.py @@ -8,7 +8,7 @@ LABEL = ["A", "B", "C", "D"] if __name__ == "__main__": parser = argparse.ArgumentParser() - parser.add_argument("--port", "-p", type=int, default=11434) + parser.add_argument("--port", "-p", type=int, default=8000) parser.add_argument("--limit", "-L", type=int, default=None) parser.add_argument("--max_tokens", "-T", type=int, default=4096) parser.add_argument("--offset", "-O", type=int, default=0) diff --git a/test/null/test_viz.py b/test/null/test_viz.py index 05a0eb8140..9a1fd425b4 100644 --- a/test/null/test_viz.py +++ b/test/null/test_viz.py @@ -470,7 +470,7 @@ class TestVizProfiler(unittest.TestCase): j = load_profile(prof) event = j['layout']['NV:SDMA:0']['events'][0] gbs = sz/(dur*1e-6)*1e-9 - self.assertEqual(event['fmt'], f"{gbs:.0f} GB/s") + self.assertEqual(event['fmt'], f"{gbs:.0f} GB/s\n{sz/1e6:.0f} MB") def test_graph(self): prof = [ProfileDeviceEvent(device='NV', tdiff=decimal.Decimal(-1000)), @@ -512,7 +512,7 @@ class TestVizProfiler(unittest.TestCase): j = load_profile(prof) sdma_events = j['layout']['NV:1:SDMA:0']['events'] gbs = sz/(dur*1e-6)*1e-9 - self.assertEqual(sdma_events[0]['fmt'], f"{gbs:.0f} GB/s") + self.assertEqual(sdma_events[0]["fmt"], f"{gbs:.0f} GB/s\n{sz/1e6:.0f} MB") def test_block_ordering(self): prof = [ProfileDeviceEvent(device='NV', tdiff=decimal.Decimal(-1000)), diff --git a/test/unit/test_gguf.py b/test/unit/test_gguf.py index 5b7f2f4de5..ee042b72e4 100644 --- a/test/unit/test_gguf.py +++ b/test/unit/test_gguf.py @@ -30,11 +30,15 @@ class TestGGUF(unittest.TestCase): def test_dequantization_q4_0(self): self._test_dequantization(GGMLQuantizationType.Q4_0) def test_dequantization_q4_1(self): self._test_dequantization(GGMLQuantizationType.Q4_1) + def test_dequantization_q5_0(self): self._test_dequantization(GGMLQuantizationType.Q5_0) + def test_dequantization_q5_1(self): self._test_dequantization(GGMLQuantizationType.Q5_1) def test_dequantization_q8_0(self): self._test_dequantization(GGMLQuantizationType.Q8_0) def test_dequantization_q4_k(self): self._test_dequantization(GGMLQuantizationType.Q4_K) def test_dequantization_q5_k(self): self._test_dequantization(GGMLQuantizationType.Q5_K) def test_dequantization_q6_k(self): self._test_dequantization(GGMLQuantizationType.Q6_K) def test_dequantization_mxfp4(self): self._test_dequantization(GGMLQuantizationType.MXFP4) + @unittest.skipUnless(is_dtype_supported(dtypes.bfloat16), "Backend must support bfloat16") + def test_dequantization_bf16(self): self._test_dequantization(GGMLQuantizationType.BF16) def test_dequantization_mxfp4_old(self): def encode(nibbles, E): packed = [(low & 0xF) | ((high & 0xF) << 4) for low, high in zip(nibbles[:16], nibbles[16:])] @@ -120,17 +124,21 @@ class TestGGUF(unittest.TestCase): class TestGGUFGEMV(unittest.TestCase): def _test_gguf_gemv(self, qtype: GGMLQuantizationType): block_size, type_size = GGML_QUANT_SIZES[qtype] - rows, cols = 8192, 2048 + rows, cols = (1024, 512) if qtype == GGMLQuantizationType.BF16 else (8192, 2048) n_blocks = rows * cols // block_size rng = np.random.default_rng(42) - # generate random quantized blocks with valid fp16 scale fields (random bytes can produce NaN scales) - q_data = rng.integers(0, 256, size=n_blocks * type_size, dtype=np.uint8).reshape(n_blocks, type_size) - scales = np.float16(rng.standard_normal(n_blocks * 4)).view(np.uint8).reshape(n_blocks, -1) - if qtype == GGMLQuantizationType.Q8_0: q_data[:, :2] = scales[:, :2] # d at offset 0 - elif qtype in (GGMLQuantizationType.Q4_K, GGMLQuantizationType.Q5_K): q_data[:, :4] = scales[:, :4] # d, dmin at offset 0 - elif qtype == GGMLQuantizationType.Q6_K: q_data[:, -2:] = scales[:, :2] # d at end - elif qtype == GGMLQuantizationType.MXFP4: q_data[:, 0] = rng.integers(120, 136, size=n_blocks, dtype=np.uint8) # constrain byte0 - q_data = q_data.flatten() + if qtype == GGMLQuantizationType.BF16: + q_data = (rng.standard_normal(rows * cols).astype(np.float32).view(np.uint32) >> 16).astype(np.uint16).view(np.uint8) + else: + # generate random quantized blocks with valid fp16 scale fields (random bytes can produce NaN scales) + q_data = rng.integers(0, 256, size=n_blocks * type_size, dtype=np.uint8).reshape(n_blocks, type_size) + scales = np.float16(rng.standard_normal(n_blocks * 4)).view(np.uint8).reshape(n_blocks, -1) + if qtype in (GGMLQuantizationType.Q5_0, GGMLQuantizationType.Q8_0): q_data[:, :2] = scales[:, :2] # d at offset 0 + elif qtype in (GGMLQuantizationType.Q5_1, GGMLQuantizationType.Q4_K, GGMLQuantizationType.Q5_K): + q_data[:, :4] = scales[:, :4] # d, m/dmin at offset 0 + elif qtype == GGMLQuantizationType.Q6_K: q_data[:, -2:] = scales[:, :2] # d at end + elif qtype == GGMLQuantizationType.MXFP4: q_data[:, 0] = rng.integers(120, 136, size=n_blocks, dtype=np.uint8) # constrain byte0 + q_data = q_data.flatten() ref = dequantize(q_data, qtype).reshape(rows, cols) # build a minimal gguf in memory: header + 1 tensor info + aligned data @@ -148,15 +156,18 @@ class TestGGUFGEMV(unittest.TestCase): x = rng.standard_normal(cols).astype(np.float32) np.testing.assert_allclose((tensors["weight"] @ Tensor(x)).numpy(), ref @ x, atol=1e-2, rtol=1e-2) - # can only expect the weights to be identical if we really support float16 (ie. not decompositions) - if is_dtype_supported(dtypes.half): np.testing.assert_equal(tensors["weight"].numpy(), ref) + if qtype == GGMLQuantizationType.BF16 or is_dtype_supported(dtypes.half): np.testing.assert_equal(tensors["weight"].numpy(), ref) assert np.isfinite(ref).all() and np.isfinite(tensors["weight"].numpy()).all(), f"{qtype.name} has NaN/Inf" def test_gguf_gemv_q8_0(self): self._test_gguf_gemv(GGMLQuantizationType.Q8_0) + def test_gguf_gemv_q5_0(self): self._test_gguf_gemv(GGMLQuantizationType.Q5_0) + def test_gguf_gemv_q5_1(self): self._test_gguf_gemv(GGMLQuantizationType.Q5_1) def test_gguf_gemv_q4_k(self): self._test_gguf_gemv(GGMLQuantizationType.Q4_K) def test_gguf_gemv_q5_k(self): self._test_gguf_gemv(GGMLQuantizationType.Q5_K) def test_gguf_gemv_q6_k(self): self._test_gguf_gemv(GGMLQuantizationType.Q6_K) def test_gguf_gemv_mxfp4(self): self._test_gguf_gemv(GGMLQuantizationType.MXFP4) + @unittest.skipUnless(is_dtype_supported(dtypes.bfloat16), "Backend must support bfloat16") + def test_gguf_gemv_bf16(self): self._test_gguf_gemv(GGMLQuantizationType.BF16) if __name__ == '__main__': unittest.main() diff --git a/tinygrad/apps/llm.py b/tinygrad/apps/llm.py index 6fe4a19201..5e22fd9b86 100644 --- a/tinygrad/apps/llm.py +++ b/tinygrad/apps/llm.py @@ -287,8 +287,7 @@ models = { "olmoe": "https://huggingface.co/allenai/OLMoE-1B-7B-0924-Instruct-GGUF/resolve/main/olmoe-1b-7b-0924-instruct-q4_k_m.gguf", } -# *** simple OpenAI compatible server on 11434 to match ollama *** -# OPENAI_BASE_URL=http://localhost:11434/v1 OPENAI_API_KEY=ollama uvx --from gpt-command-line gpt +# *** simple OpenAI API compatible server with web interface on http://localhost:8000/ *** CHAT_HTML = b'''tinygrad chat