diff --git a/examples/mlperf/models/flat_llama.py b/examples/mlperf/models/flat_llama.py index 018d61f7aa..66a982b410 100644 --- a/examples/mlperf/models/flat_llama.py +++ b/examples/mlperf/models/flat_llama.py @@ -2,9 +2,8 @@ import math, os if __name__ == "__main__": os.environ["DEFAULT_FLOAT"] = "bfloat16" os.environ["OPTIM_DTYPE"] = "bfloat16" - if "DEV" not in os.environ: os.environ["DEV"] = "NULL" + if "DEV" not in os.environ: os.environ["DEV"] = "NULL::gfx950" # CDNA - os.environ["EMULATE"] = "AMD_CDNA4" os.environ["DEVICE_IN_FUNCTION_BUG"] = "1" os.environ["ALL2ALL"] = "1" os.environ["USE_ATOMICS"] = "1" @@ -23,7 +22,7 @@ ASM_GEMM = getenv("ASM_GEMM", 0) FUSED_INPUT_QUANTIZE = getenv("FUSED_INPUT_QUANTIZE", 0) FUSED_ADD_NORM_MUL_QUANTIZE = getenv("FUSED_ADD_NORM_MUL_QUANTIZE", 0) FUSED_SILU_W13 = getenv("FUSED_SILU_W13", 0) -SPLIT_W13 = getenv("SPLIT_W13", 0) +SPLIT_W13 = getenv("SPLIT_W13", 1) FP8_DTYPE = dtypes.fp8e4m3 FP8_GRAD_DTYPE = dtypes.fp8e5m2 @@ -212,8 +211,9 @@ class FlatTransformer: attn, attn_amaxs, attn_saves = self.attention(x, freqs_cis, **attn_kwargs) ffn, h, ffn_amaxs, ffn_saves = self.feed_forward(x, attn, **ffn_kwargs) h = h + ffn - if save: return (h, *attn_amaxs, *ffn_amaxs, *attn_saves, *ffn_saves) - else: return (h, *attn_amaxs, *ffn_amaxs) + amaxs = tuple(a.detach() for a in (*attn_amaxs, *ffn_amaxs)) + if save: return (h, *amaxs, *attn_saves, *ffn_saves) + else: return (h, *amaxs) def shard(self, device:tuple[str, ...], mp:bool=False): from tinygrad.nn.state import get_parameters @@ -292,7 +292,7 @@ if __name__ == "__main__": SEQLEN = config["SEQLEN"] = getenv("SEQLEN", 8192) from examples.llama3 import MODEL_PARAMS - model_params = MODEL_PARAMS[getenv("LLAMA3_SIZE", "8B")]["args"] + model_params = MODEL_PARAMS[llama_size:=getenv("LLAMA3_SIZE", "8B")]["args"] if (llama_layers:=getenv("LLAMA_LAYERS")) != 0: model_params['n_layers'] = llama_layers model = FlatTransformer(**model_params, max_context=SEQLEN) state = nn.state.get_state_dict(model) @@ -306,8 +306,12 @@ if __name__ == "__main__": model.shard(tuple(f"{Device.DEFAULT}:{i}" for i in range(MP)), mp=True) # preallocate all the grad buffers and zero them out - grads = {x:Tensor.zeros(x.shape, dtype=x.dtype, device=x.device).contiguous() - for x in state.values() if x.requires_grad} + grad_dtype = lambda x: dtypes.bfloat16 if x.dtype in dtypes.fp8s else x.dtype + def _make_grad(x): + if isinstance(x.device, tuple) and x.uop.axis is not None: + return Tensor.zeros(x.shape, dtype=grad_dtype(x), device=x.device[0]).shard_(x.device, axis=x.uop.axis).contiguous() + return Tensor.zeros(x.shape, dtype=grad_dtype(x), device=x.device).contiguous() + grads = {x:_make_grad(x) for x in state.values() if x.requires_grad} # print model size sz = 0 @@ -323,16 +327,22 @@ if __name__ == "__main__": if MP > 1: tokens = tokens.shard(tuple(f"{Device.DEFAULT}:{i}" for i in range(MP))) @TinyJit - def jit_step(tokens:Tensor): - with Timing("python forward: "): loss = model(tokens[:, :-1]).sparse_categorical_crossentropy(tokens[:, 1:]) + def fwd_bwd(tokens:Tensor): + with Timing("python forward: "): loss = model(tokens[:, :-1], save=llama_size=="8B").sparse_categorical_crossentropy(tokens[:, 1:]) with Timing("python backward: "): for t,g in zip(grads, loss.gradient(*grads)): apply_grad(grads[t], g.uop) - with Timing("run step: "): loss.realize(*grads.values()) + with Timing("run fwd_bwd: "): loss.realize(*grads.values()) + + @TinyJit + def optim_step(): + for g in grads.values(): g.assign(g.zeros_like()) + Tensor.realize(*grads.values()) for i in range(6): GlobalCounters.reset() profile_marker(f"step {i}") with Timing(colored(f"*** step {i}: ", "red")): - jit_step(tokens) + fwd_bwd(tokens) + optim_step() print("mem per device: " + ', '.join(f"{dev}: {mem/1e9:.2f} GB" for dev, mem in sorted(GlobalCounters.mem_used_per_device.items())))