Gpt2
gpt2.ml
open Kaun
module Hf = Kaun_hf
let invalid_argf fmt = Printf.ksprintf invalid_arg fmt
type config = {
vocab_size : int;
n_positions : int;
n_embd : int;
n_layer : int;
n_head : int;
n_inner : int;
layer_norm_eps : float;
}
type 'b block = {
ln1 : 'b Layer_norm.params;
attn : 'b Attention.params;
ln2 : 'b Layer_norm.params;
fc : 'b Linear.params;
proj : 'b Linear.params;
}
type 'b params = {
wte : 'b Embedding.params;
wpe : 'b Embedding.params;
blocks : 'b block list;
ln_f : 'b Layer_norm.params;
}
type t = Nx.float32_elt params
module Params = struct
type nonrec t = t
let map_block (f : 'a 'b. ('a, 'b) Nx.t -> ('a, 'b) Nx.t) b =
{
ln1 = Layer_norm.map f b.ln1;
attn = Attention.map f b.attn;
ln2 = Layer_norm.map f b.ln2;
fc = Linear.map f b.fc;
proj = Linear.map f b.proj;
}
let map (f : 'a 'b. ('a, 'b) Nx.t -> ('a, 'b) Nx.t) p =
{
wte = Embedding.map f p.wte;
wpe = Embedding.map f p.wpe;
blocks = List.map (map_block f) p.blocks;
ln_f = Layer_norm.map f p.ln_f;
}
let map2_block (f : 'a 'b. ('a, 'b) Nx.t -> ('a, 'b) Nx.t -> ('a, 'b) Nx.t) b
b' =
{
ln1 = Layer_norm.map2 f b.ln1 b'.ln1;
attn = Attention.map2 f b.attn b'.attn;
ln2 = Layer_norm.map2 f b.ln2 b'.ln2;
fc = Linear.map2 f b.fc b'.fc;
proj = Linear.map2 f b.proj b'.proj;
}
let map2 (f : 'a 'b. ('a, 'b) Nx.t -> ('a, 'b) Nx.t -> ('a, 'b) Nx.t) p p' =
{
wte = Embedding.map2 f p.wte p'.wte;
wpe = Embedding.map2 f p.wpe p'.wpe;
blocks = List.map2 (map2_block f) p.blocks p'.blocks;
ln_f = Layer_norm.map2 f p.ln_f p'.ln_f;
}
let iter_block (f : 'a 'b. ('a, 'b) Nx.t -> unit) b =
Layer_norm.iter f b.ln1;
Attention.iter f b.attn;
Layer_norm.iter f b.ln2;
Linear.iter f b.fc;
Linear.iter f b.proj
let iter (f : 'a 'b. ('a, 'b) Nx.t -> unit) p =
Embedding.iter f p.wte;
Embedding.iter f p.wpe;
List.iter (iter_block f) p.blocks;
Layer_norm.iter f p.ln_f
let names p =
let pre prefix ns = List.map (fun n -> prefix ^ "." ^ n) ns in
let block b =
pre "ln1" (Layer_norm.names b.ln1)
@ pre "attn" (Attention.names b.attn)
@ pre "ln2" (Layer_norm.names b.ln2)
@ pre "fc" (Linear.names b.fc)
@ pre "proj" (Linear.names b.proj)
in
pre "wte" (Embedding.names p.wte)
@ pre "wpe" (Embedding.names p.wpe)
@ List.concat
(List.mapi
(fun i b -> pre (Printf.sprintf "blocks.%d" i) (block b))
p.blocks)
@ pre "ln_f" (Layer_norm.names p.ln_f)
end
let astype dt p =
let block b =
{
ln1 = Layer_norm.astype dt b.ln1;
attn = Attention.astype dt b.attn;
ln2 = Layer_norm.astype dt b.ln2;
fc = Linear.astype dt b.fc;
proj = Linear.astype dt b.proj;
}
in
{
wte = Embedding.astype dt p.wte;
wpe = Embedding.astype dt p.wpe;
blocks = List.map block p.blocks;
ln_f = Layer_norm.astype dt p.ln_f;
}
let make cfg =
let zeros = Init.zeros in
let linear ~inputs ~outputs =
Linear.make ~w_init:zeros ~bias_init:zeros ~inputs ~outputs Nx.float32
in
let embedding ~vocab = Embedding.make ~init:zeros ~vocab ~dim:cfg.n_embd in
let block () =
{
ln1 = Layer_norm.init ~dim:cfg.n_embd;
attn =
Attention.make ~w_init:zeros ~bias_init:zeros ~embed_dim:cfg.n_embd
Nx.float32;
ln2 = Layer_norm.init ~dim:cfg.n_embd;
fc = linear ~inputs:cfg.n_embd ~outputs:cfg.n_inner;
proj = linear ~inputs:cfg.n_inner ~outputs:cfg.n_embd;
}
in
{
wte = embedding ~vocab:cfg.vocab_size Nx.float32;
wpe = embedding ~vocab:cfg.n_positions Nx.float32;
blocks = List.init cfg.n_layer (fun _ -> block ());
ln_f = Layer_norm.init ~dim:cfg.n_embd;
}
let drop dropout i x =
match dropout with
| None -> x
| Some (rate, key) ->
Dropout.apply ~rate ~training:true ~key:(Rune.Rng.fold_in key i) x
let block_apply cfg ?dropout b x =
let eps = cfg.layer_norm_eps in
let x =
Nx.add x
(drop dropout 0
(Attention.apply ~num_heads:cfg.n_head ~causal:true b.attn
(Layer_norm.apply ~eps b.ln1 x)))
in
Nx.add x
(drop dropout 1
(Linear.apply b.proj
(Fn.gelu_approx (Linear.apply b.fc (Layer_norm.apply ~eps b.ln2 x)))))
let logits cfg ?dropout p ids =
let seq = (Nx.shape ids).(1) in
if seq > cfg.n_positions then
invalid_argf "Gpt2.logits: seq %d exceeds n_positions %d" seq
cfg.n_positions;
let sub i =
Option.map (fun (rate, key) -> (rate, Rune.Rng.fold_in key i)) dropout
in
let pos = Nx.reshape [| 1; seq |] (Nx.arange Nx.int32 0 seq 1) in
let x =
drop dropout 0
(Nx.add (Embedding.apply p.wte ids) (Embedding.apply p.wpe pos))
in
let _, x =
List.fold_left
(fun (i, x) b -> (i + 1, block_apply cfg ?dropout:(sub i) b x))
(1, x) p.blocks
in
let h = Layer_norm.apply ~eps:cfg.layer_norm_eps p.ln_f x in
Nx.matmul h (Nx.transpose p.wte.table)
type 'b cache = 'b Attention.Cache.t list
let cache cfg ~len dtype =
List.init cfg.n_layer (fun _ ->
Attention.Cache.make ~num_heads:cfg.n_head
~head_dim:(cfg.n_embd / cfg.n_head) ~len dtype)
let block_apply_cached cfg b c ~pos x =
let eps = cfg.layer_norm_eps in
let attn, c =
Attention.apply_cached ~num_heads:cfg.n_head ~pos ~cache:c b.attn
(Layer_norm.apply ~eps b.ln1 x)
in
let x = Nx.add x attn in
( Nx.add x
(Linear.apply b.proj
(Fn.gelu_approx (Linear.apply b.fc (Layer_norm.apply ~eps b.ln2 x)))),
c )
let logits_cached cfg p ~pos caches ids =
let seq = (Nx.shape ids).(1) in
let positions =
Nx.add
(Nx.reshape [| 1; 1 |] pos)
(Nx.reshape [| 1; seq |] (Nx.arange Nx.int32 0 seq 1))
in
let x =
Nx.add (Embedding.apply p.wte ids) (Embedding.apply p.wpe positions)
in
let x, rev_caches =
List.fold_left2
(fun (x, cs) b c ->
let x, c = block_apply_cached cfg b c ~pos x in
(x, c :: cs))
(x, []) p.blocks caches
in
let x = Nx.slice [ A; R (seq - 1, seq) ] x in
let h = Layer_norm.apply ~eps:cfg.layer_norm_eps p.ln_f x in
let vocab = (Nx.shape p.wte.table).(0) in
( Nx.reshape
[| (Nx.shape ids).(0); vocab |]
(Nx.matmul h (Nx.transpose p.wte.table)),
List.rev rev_caches )
let generate cfg p ~max_tokens prompt =
let n0 = Array.length prompt in
if n0 = 0 then invalid_arg "Gpt2.generate: prompt must not be empty";
let tokens = Array.make (n0 + max_tokens) 0l in
Array.blit prompt 0 tokens 0 n0;
for n = n0 to n0 + max_tokens - 1 do
let input = Nx.create Nx.int32 [| 1; n |] (Array.sub tokens 0 n) in
let last = Nx.slice [ I 0; I (n - 1) ] (logits cfg p input) in
tokens.(n) <- Nx.item [] (Nx.argmax ~axis:0 last)
done;
tokens
let hf_name name =
match name with
| "wte.weight" -> "wte.table"
| "wpe.weight" -> "wpe.table"
| "ln_f.weight" -> "ln_f.gamma"
| "ln_f.bias" -> "ln_f.beta"
| _ -> (
match String.split_on_char '.' name with
| "h" :: i :: rest -> (
let ours leaf = Printf.sprintf "blocks.%s.%s" i leaf in
match rest with
| [ "ln_1"; "weight" ] -> ours "ln1.gamma"
| [ "ln_1"; "bias" ] -> ours "ln1.beta"
| [ "ln_2"; "weight" ] -> ours "ln2.gamma"
| [ "ln_2"; "bias" ] -> ours "ln2.beta"
| [ "attn"; "c_proj"; "weight" ] -> ours "attn.out.w"
| [ "attn"; "c_proj"; "bias" ] -> ours "attn.out.b"
| [ "mlp"; "c_fc"; "weight" ] -> ours "fc.w"
| [ "mlp"; "c_fc"; "bias" ] -> ours "fc.b"
| [ "mlp"; "c_proj"; "weight" ] -> ours "proj.w"
| [ "mlp"; "c_proj"; "bias" ] -> ours "proj.b"
| _ -> name )
| _ -> name)
let of_hf ~n_layer ckpt =
let split_qkv ckpt i =
let fused leaf = Printf.sprintf "h.%d.attn.c_attn.%s" i leaf in
let ours p leaf = Printf.sprintf "blocks.%d.attn.%s.%s" i p leaf in
ckpt
|> Hf.split (fused "weight")
~into:[ ours "q" "w"; ours "k" "w"; ours "v" "w" ]
|> Hf.split (fused "bias")
~into:[ ours "q" "b"; ours "k" "b"; ours "v" "b" ]
in
List.fold_left split_qkv ckpt (List.init n_layer Fun.id) |> Hf.rename hf_name
let json_mem name = function
| Jsont.Object (mems, _) -> (
match Jsont.Json.find_mem name mems with
| Some (_, v) -> v
| None -> Jsont.Null ((), Jsont.Meta.none))
| _ -> Jsont.Null ((), Jsont.Meta.none)
let json_int ~default json name =
match json_mem name json with
| Jsont.Number (f, _) -> int_of_float f
| _ -> default ()
let config_of_json json =
let req name =
json_int
~default:(fun () -> failwith ("gpt2 config.json: missing " ^ name))
json name
in
let n_embd = req "n_embd" in
{
vocab_size = req "vocab_size";
n_positions = json_int ~default:(fun () -> 1024) json "n_positions";
n_embd;
n_layer = req "n_layer";
n_head = req "n_head";
n_inner = json_int ~default:(fun () -> 4 * n_embd) json "n_inner";
layer_norm_eps =
(match json_mem "layer_norm_epsilon" json with
| Jsont.Number (f, _) -> f
| _ -> 1e-5);
}
let of_checkpoint cfg ckpt =
of_hf ~n_layer:cfg.n_layer ckpt
|> Checkpoint.to_params (module Params) ~like:(make cfg) ~cast:true
let from_file cfg path = of_checkpoint cfg (Checkpoint.load path)
let from_pretrained ?(repo_id = "gpt2") () =
let cfg = config_of_json (Hf.load_config repo_id) in
(cfg, of_checkpoint cfg (Hf.load_checkpoint repo_id))
main.ml
let default_prompt = "What is the answer to life, the universe, and everything?"
let local_file file =
let cache =
try Sys.getenv "XDG_CACHE_HOME"
with Not_found -> Filename.concat (Sys.getenv "HOME") ".cache"
in
let path = List.fold_left Filename.concat cache [ "tolk-gpt2"; file ] in
if Sys.file_exists path then Some path else None
let gpt2_124m : Gpt2.config =
{
vocab_size = 50257;
n_positions = 1024;
n_embd = 768;
n_layer = 12;
n_head = 12;
n_inner = 3072;
layer_norm_eps = 1e-5;
}
let load_model () =
match local_file "model.safetensors" with
| Some path -> (gpt2_124m, Gpt2.from_file gpt2_124m path)
| None -> Gpt2.from_pretrained ()
let load_tokenizer () =
let path =
match local_file "tokenizer.json" with
| Some path -> path
| None -> Kaun_hf.download_file ~file:"tokenizer.json" "gpt2"
in
match Brot.from_file path with
| Ok t -> t
| Error e -> failwith ("tokenizer: " ^ e)
let generate_jit (type b) ~device cfg (params : b Gpt2.params)
(dt : (float, b) Nx.dtype) ~max_tokens prompt =
let module Step = struct
type t = {
token : Nx.int32_t;
pos : Nx.int32_t;
caches : b Gpt2.cache;
}
let map (f : 'a 'c. ('a, 'c) Nx.t -> ('a, 'c) Nx.t) { token; pos; caches } =
{
token = f token;
pos = f pos;
caches = List.map (Kaun.Attention.Cache.map f) caches;
}
let map2 (f : 'a 'c. ('a, 'c) Nx.t -> ('a, 'c) Nx.t -> ('a, 'c) Nx.t) a b =
{
token = f a.token b.token;
pos = f a.pos b.pos;
caches = List.map2 (Kaun.Attention.Cache.map2 f) a.caches b.caches;
}
let iter (f : 'a 'c. ('a, 'c) Nx.t -> unit) { token; pos; caches } =
f token;
f pos;
List.iter (Kaun.Attention.Cache.iter f) caches
end in
let n0 = Array.length prompt in
let len = n0 + max_tokens in
let tokens = Array.make len 0l in
Array.blit prompt 0 tokens 0 n0;
let step_fn =
Rune.jit2 ~device
(module Step)
(module Step)
(fun { Step.token; pos; caches } ->
let seq = (Nx.shape token).(1) in
let logits, caches = Gpt2.logits_cached cfg params ~pos caches token in
{
Step.token = Nx.reshape [| 1; 1 |] (Nx.argmax ~axis:1 logits);
pos = Nx.add_s pos (Int32.of_int seq);
caches;
})
in
let t0 = Unix.gettimeofday () in
let state =
ref
(step_fn
{
Step.token = Nx.create Nx.int32 [| 1; n0 |] prompt;
pos = Nx.zeros Nx.int32 [| 1 |];
caches = Gpt2.cache cfg ~len dt;
})
in
tokens.(n0) <- Nx.item [ 0; 0 ] !state.Step.token;
let prefill = Unix.gettimeofday () -. t0 in
let times = Array.make (max 1 (max_tokens - 1)) 0. in
for n = n0 + 1 to len - 1 do
let t0 = Unix.gettimeofday () in
let s = step_fn !state in
tokens.(n) <- Nx.item [ 0; 0 ] s.Step.token;
state := s;
times.(n - n0 - 1) <- Unix.gettimeofday () -. t0
done;
if max_tokens > 2 then begin
let rest = Array.fold_left ( +. ) (-.times.(0)) times in
Printf.printf
"prefill %.2f s (compile), first step %.2f s (compile), then %.2f tok/s\n\
%!"
prefill times.(0)
(float_of_int (max_tokens - 2) /. rest)
end;
tokens
let () =
let prompt = ref default_prompt in
let count = ref 10 in
let jit = ref "" in
let dtype = ref "float32" in
Arg.parse
[
("--prompt", Arg.Set_string prompt, "Phrase to start with");
("--count", Arg.Set_int count, "Max number of tokens to generate");
( "--jit",
Arg.Set_string jit,
"Compile the forward pass for this device (CPU or CUDA); eager when \
omitted" );
( "--dtype",
Arg.Set_string dtype,
"Model dtype: float32 (default), float16 or bfloat16. Half dtypes cast \
the float32 checkpoint once and run weights, activations and caches \
at half precision" );
]
(fun a -> raise (Arg.Bad ("unexpected argument " ^ a)))
"gpt2 [--prompt P] [--count N] [--jit DEVICE] [--dtype DT]";
let tokenizer = load_tokenizer () in
let t0 = Unix.gettimeofday () in
let cfg, params = load_model () in
Printf.printf "loaded weights in %.2f s\n%!" (Unix.gettimeofday () -. t0);
let ids = Array.map Int32.of_int (Brot.encode_ids tokenizer !prompt) in
let run : type b. (float, b) Nx.dtype -> b Gpt2.params -> int32 array =
fun dt params ->
let bytes = ref 0 in
let count_t t = bytes := !bytes + Nx.nbytes t in
Kaun.Embedding.iter count_t params.wte;
Kaun.Embedding.iter count_t params.wpe;
List.iter
(fun (b : b Gpt2.block) ->
Kaun.Layer_norm.iter count_t b.ln1;
Kaun.Attention.iter count_t b.attn;
Kaun.Layer_norm.iter count_t b.ln2;
Kaun.Linear.iter count_t b.fc;
Kaun.Linear.iter count_t b.proj)
params.blocks;
Kaun.Layer_norm.iter count_t params.ln_f;
Printf.printf "weights: %.0f MB at %s\n%!"
(float_of_int !bytes /. 1e6)
!dtype;
match !jit with
| "" -> Gpt2.generate cfg params ~max_tokens:!count ids
| device -> generate_jit ~device cfg params dt ~max_tokens:!count ids
in
let t0 = Unix.gettimeofday () in
let toks =
match !dtype with
| "float32" -> run Nx.float32 params
| "float16" -> run Nx.float16 (Gpt2.astype Nx.float16 params)
| "bfloat16" -> run Nx.bfloat16 (Gpt2.astype Nx.bfloat16 params)
| d -> failwith ("--dtype must be float32, float16 or bfloat16, got " ^ d)
in
let dt = Unix.gettimeofday () -. t0 in
Printf.printf "generated %d tokens in %.2f s (%.2f tok/s)\n%!" !count dt
(float_of_int !count /. dt);
let text = Brot.decode tokenizer (Array.map Int32.to_int toks) in
print_endline text
train.ml
open Kaun
let gpt2_124m : Gpt2.config =
{
vocab_size = 50257;
n_positions = 1024;
n_embd = 768;
n_layer = 12;
n_head = 12;
n_inner = 3072;
layer_norm_eps = 1e-5;
}
let batch_size = 4
let seq_len = 64
let lr = 1e-4
let repr_float x =
let rec go p =
let s = Printf.sprintf "%.*g" p x in
if p >= 17 || float_of_string s = x then s else go (p + 1)
in
go 1
let read_file path =
let ic = open_in_bin path in
Fun.protect
~finally:(fun () -> close_in ic)
(fun () -> really_input_string ic (in_channel_length ic))
let load_ids path =
let json =
match Jsont_bytesrw.decode_string Jsont.json (read_file path) with
| Ok v -> v
| Error e -> failwith (path ^ ": " ^ e)
in
let ids =
match json with
| Jsont.Object (mems, _) -> (
match Jsont.Json.find_mem "ids" mems with
| Some (_, ids) -> ids
| None -> failwith (path ^ ": no \"ids\" member"))
| _ -> failwith (path ^ ": expected an object")
in
let row = function
| Jsont.Array (cells, _) ->
Array.of_list
(List.map
(function
| Jsont.Number (f, _) -> Int32.of_float f
| _ -> failwith (path ^ ": non-numeric id"))
cells)
| _ -> failwith (path ^ ": expected rows of ids")
in
match ids with
| Jsont.Array (rows, _) -> Array.of_list (List.map row rows)
| _ -> failwith (path ^ ": expected an array of rows")
let batch_of_ids ids =
if
Array.length ids <> batch_size
|| Array.exists (fun r -> Array.length r <> seq_len + 1) ids
then
failwith
(Printf.sprintf "token grid must be %d x %d" batch_size (seq_len + 1));
let take off =
Array.init (batch_size * seq_len) (fun i ->
ids.(i / seq_len).(off + (i mod seq_len)))
in
( Nx.create Nx.int32 [| batch_size; seq_len |] (take 0),
Nx.create Nx.int32 [| batch_size * seq_len |] (take 1) )
let loss_fn inputs targets ?dropout params =
let logits = Gpt2.logits gpt2_124m ?dropout params inputs in
let logits =
Nx.reshape [| batch_size * seq_len; gpt2_124m.vocab_size |] logits
in
Loss.softmax_cross_entropy_sparse logits targets
let loss_fn_half compute inputs targets ?dropout params =
let params = Gpt2.astype compute params in
let logits = Gpt2.logits gpt2_124m ?dropout params inputs in
let logits =
Nx.reshape [| batch_size * seq_len; gpt2_124m.vocab_size |] logits
in
Loss.softmax_cross_entropy_sparse (Nx.cast Nx.float32 logits) targets
module Step_out = struct
type t = { params : Gpt2.t; loss : Nx.float32_t }
let map (f : 'a 'b. ('a, 'b) Nx.t -> ('a, 'b) Nx.t) t =
{ params = Gpt2.Params.map f t.params; loss = f t.loss }
let map2 (f : 'a 'b. ('a, 'b) Nx.t -> ('a, 'b) Nx.t -> ('a, 'b) Nx.t) a b =
{ params = Gpt2.Params.map2 f a.params b.params; loss = f a.loss b.loss }
let iter (f : 'a 'b. ('a, 'b) Nx.t -> unit) t =
Gpt2.Params.iter f t.params;
f t.loss
end
let train_step objective params =
let loss, grads = Rune.value_and_grad (module Gpt2.Params) objective params in
let state = Vega.sgd_init (module Gpt2.Params) params in
let params =
fst (Vega.sgd_step (module Gpt2.Params) ~lr state ~params ~grads)
in
{ Step_out.params; loss }
let map2_key f a b =
match (a, b) with
| Some a, Some b -> Some (f a b)
| None, None -> None
| _ -> invalid_arg "dropout key: presence mismatch"
module Keyed_in = struct
type t = { params : Gpt2.t; key : Rune.Rng.key option }
let map (f : 'a 'b. ('a, 'b) Nx.t -> ('a, 'b) Nx.t) t =
{ params = Gpt2.Params.map f t.params; key = Option.map f t.key }
let map2 (f : 'a 'b. ('a, 'b) Nx.t -> ('a, 'b) Nx.t -> ('a, 'b) Nx.t) a b =
{
params = Gpt2.Params.map2 f a.params b.params;
key = map2_key f a.key b.key;
}
let iter (f : 'a 'b. ('a, 'b) Nx.t -> unit) t =
Gpt2.Params.iter f t.params;
Option.iter f t.key
end
module Scaled_in = struct
type t = {
params : Gpt2.t;
ls : Vega.Loss_scale.t;
key : Rune.Rng.key option;
}
let map (f : 'a 'b. ('a, 'b) Nx.t -> ('a, 'b) Nx.t) t =
{
params = Gpt2.Params.map f t.params;
ls = Vega.Loss_scale.map f t.ls;
key = Option.map f t.key;
}
let map2 (f : 'a 'b. ('a, 'b) Nx.t -> ('a, 'b) Nx.t -> ('a, 'b) Nx.t) a b =
{
params = Gpt2.Params.map2 f a.params b.params;
ls = Vega.Loss_scale.map2 f a.ls b.ls;
key = map2_key f a.key b.key;
}
let iter (f : 'a 'b. ('a, 'b) Nx.t -> unit) t =
Gpt2.Params.iter f t.params;
Vega.Loss_scale.iter f t.ls;
Option.iter f t.key
end
module Scaled_out = struct
type t = { params : Gpt2.t; loss : Nx.float32_t; ls : Vega.Loss_scale.t }
let map (f : 'a 'b. ('a, 'b) Nx.t -> ('a, 'b) Nx.t) t =
{
params = Gpt2.Params.map f t.params;
loss = f t.loss;
ls = Vega.Loss_scale.map f t.ls;
}
let map2 (f : 'a 'b. ('a, 'b) Nx.t -> ('a, 'b) Nx.t -> ('a, 'b) Nx.t) a b =
{
params = Gpt2.Params.map2 f a.params b.params;
loss = f a.loss b.loss;
ls = Vega.Loss_scale.map2 f a.ls b.ls;
}
let iter (f : 'a 'b. ('a, 'b) Nx.t -> unit) t =
Gpt2.Params.iter f t.params;
f t.loss;
Vega.Loss_scale.iter f t.ls
end
let train_step_scaled objective { Scaled_in.params; ls; key } =
let loss, grads =
Rune.vjp (module Gpt2.Params) (objective key) params
ls.Vega.Loss_scale.scale
in
let grads = Vega.Loss_scale.unscale (module Gpt2.Params) ls grads in
let finite = Vega.Loss_scale.grads_finite (module Gpt2.Params) grads in
let state = Vega.sgd_init (module Gpt2.Params) params in
let params' =
fst (Vega.sgd_step (module Gpt2.Params) ~lr state ~params ~grads)
in
let params =
Gpt2.Params.map2 (fun p p' -> Nx.where finite p' p) params params'
in
{ Scaled_out.params; loss; ls = Vega.Loss_scale.adjust ls ~finite }
module Step_in = struct
type t = {
params : Gpt2.t;
inputs : (int32, Nx.int32_elt) Nx.t;
targets : (int32, Nx.int32_elt) Nx.t;
key : Rune.Rng.key option;
}
let map (f : 'a 'b. ('a, 'b) Nx.t -> ('a, 'b) Nx.t) t =
{
params = Gpt2.Params.map f t.params;
inputs = f t.inputs;
targets = f t.targets;
key = Option.map f t.key;
}
let map2 (f : 'a 'b. ('a, 'b) Nx.t -> ('a, 'b) Nx.t -> ('a, 'b) Nx.t) a b =
{
params = Gpt2.Params.map2 f a.params b.params;
inputs = f a.inputs b.inputs;
targets = f a.targets b.targets;
key = map2_key f a.key b.key;
}
let iter (f : 'a 'b. ('a, 'b) Nx.t -> unit) t =
Gpt2.Params.iter f t.params;
f t.inputs;
f t.targets;
Option.iter f t.key
end
let parse_devices s =
match int_of_string_opt s with
| Some n when n > 0 -> List.init n (fun i -> Printf.sprintf "CPU:%d" (i + 1))
| Some _ -> failwith "--devices: the device count must be positive"
| None -> List.map String.trim (String.split_on_char ',' s)
let rec pairwise ~abs a lo n =
if n <= 128 then begin
let s = ref 0.0 in
for i = lo to lo + n - 1 do
let x = Bigarray.Array1.unsafe_get a i in
s := !s +. if abs then Float.abs x else x
done;
!s
end
else
let h = n / 2 in
pairwise ~abs a lo h +. pairwise ~abs a (lo + h) (n - h)
let fingerprint t =
let n = Nx.numel t in
let a =
Bigarray.array1_of_genarray
(Bigarray.reshape (Nx.to_bigarray (Nx.contiguous t)) [| n |])
in
let first8 =
List.init (min 8 n) (fun i -> repr_float (Bigarray.Array1.get a i))
in
( repr_float (pairwise ~abs:false a 0 n),
repr_float (pairwise ~abs:true a 0 n),
first8 )
let fingerprint_tensors (p : Gpt2.t) =
let b0 = List.nth p.blocks 0 and b11 = List.nth p.blocks 11 in
let c_attn =
Nx.concatenate ~axis:1 [ b0.attn.q.w; b0.attn.k.w; b0.attn.v.w ]
in
[
("wte.weight", p.wte.table);
("h.0.attn.c_attn.weight", Nx.transpose c_attn);
("h.0.mlp.c_fc.weight", Nx.transpose b0.fc.w);
("h.11.attn.c_proj.weight", Nx.transpose b11.attn.out.w);
("ln_f.weight", p.ln_f.gamma);
("wpe.weight", p.wpe.table);
]
let bias name = function
| Some b -> b
| None -> failwith (name ^ ": expected a bias")
let checkpoint_of_params original (p : Gpt2.t) =
let tensor = Checkpoint.of_tensor in
let block i (b : Nx.float32_elt Gpt2.block) =
let key leaf = Printf.sprintf "h.%d.%s" i leaf in
let linear leaf (l : Linear.t) =
[
tensor (key (leaf ^ ".weight")) l.w;
tensor (key (leaf ^ ".bias")) (bias (key leaf) l.b);
]
in
[
tensor (key "ln_1.weight") b.ln1.gamma;
tensor (key "ln_1.bias") b.ln1.beta;
tensor (key "attn.c_attn.weight")
(Nx.concatenate ~axis:1 [ b.attn.q.w; b.attn.k.w; b.attn.v.w ]);
tensor (key "attn.c_attn.bias")
(Nx.concatenate
[
bias (key "attn.q") b.attn.q.b;
bias (key "attn.k") b.attn.k.b;
bias (key "attn.v") b.attn.v.b;
]);
(match Checkpoint.find (key "attn.bias") original with
| Some (Rune.Ptree.P mask) -> tensor (key "attn.bias") mask
| None -> failwith (key "attn.bias" ^ ": missing from the input file"));
tensor (key "ln_2.weight") b.ln2.gamma;
tensor (key "ln_2.bias") b.ln2.beta;
]
@ linear "attn.c_proj" b.attn.out
@ linear "mlp.c_fc" b.fc @ linear "mlp.c_proj" b.proj
in
Checkpoint.concat
([
tensor "wte.weight" p.wte.table;
tensor "wpe.weight" p.wpe.table;
tensor "ln_f.weight" p.ln_f.gamma;
tensor "ln_f.bias" p.ln_f.beta;
]
@ List.concat (List.mapi block p.blocks))
let emit_metrics buf ~device ~steps ~n_params ~initial_loss records =
let str s = "\"" ^ s ^ "\"" in
Buffer.add_string buf
(Printf.sprintf
"{\n\
\ \"device\": %s,\n\
\ \"steps\": %d,\n\
\ \"batch_size\": %d,\n\
\ \"seq_len\": %d,\n\
\ \"lr\": %s,\n\
\ \"n_params\": %d,\n\
\ \"initial_loss\": %s,\n\
\ \"jit_calls\": null,\n\
\ \"records\": [\n"
(str device) steps batch_size seq_len (repr_float lr) n_params
(str initial_loss));
List.iteri
(fun i (step, loss, ms, weights) ->
if i > 0 then Buffer.add_string buf ",\n";
Buffer.add_string buf
(Printf.sprintf
" {\"step\": %d, \"loss\": %s, \"ms\": %s,\n \"weights\": {" step
(str loss) (repr_float ms));
List.iteri
(fun j (name, (sum, abs_sum, first8)) ->
if j > 0 then Buffer.add_string buf ",";
Buffer.add_string buf
(Printf.sprintf
"\n %s: {\"sum\": %s, \"abs_sum\": %s, \"first8\": [%s]}"
(str name) (str sum) (str abs_sum)
(String.concat ", " (List.map str first8))))
weights;
Buffer.add_string buf "}}")
records;
Buffer.add_string buf "\n ]\n}\n"
let () =
let device = ref "CPU" in
let devices = ref "" in
let compute_dtype = ref "float32" in
let steps = ref 20 in
let dropout = ref 0.0 in
let seed = ref 0 in
let here = Filename.dirname Sys.executable_name in
let tokens = ref (Filename.concat here "tokens.json") in
let model =
ref
(List.fold_left Filename.concat
(try Sys.getenv "HOME" with Not_found -> ".")
[ ".cache"; "tolk-gpt2"; "model.safetensors" ])
in
let metrics_out = ref "" in
let save_weights = ref "" in
Arg.parse
[
("--device", Arg.Set_string device, "Device to jit for (CPU or CUDA)");
( "--devices",
Arg.Set_string devices,
"Data-parallel device tuple: a CPU count (2 = CPU:1,CPU:2) or a \
comma-separated list (CUDA:0,CUDA:1)" );
( "--compute-dtype",
Arg.Set_string compute_dtype,
"Forward/backward dtype: float32 (default), bfloat16 or float16. \
Parameters, gradients and the optimizer step stay float32 (master \
weights); float16 adds dynamic loss scaling" );
("--steps", Arg.Set_int steps, "Number of training steps");
( "--dropout",
Arg.Set_float dropout,
"Dropout rate at the GPT-2 sites (default 0: the reference protocol's \
dropout-free graph)" );
( "--seed",
Arg.Set_int seed,
"Root PRNG seed for the per-step dropout keys (default 0)" );
("--tokens", Arg.Set_string tokens, "Path to the tokens.json grid");
("--model", Arg.Set_string model, "Path to the initial safetensors");
("--metrics-out", Arg.Set_string metrics_out, "Metrics JSON path");
( "--save-weights",
Arg.Set_string save_weights,
"Save final weights (input-file layout) to this path" );
]
(fun a -> raise (Arg.Bad ("unexpected argument " ^ a)))
"train --device DEV [--devices TUPLE] [--steps N] [--dropout R] [--seed N] \
[--tokens F] [--model F] [--metrics-out F] [--save-weights F]";
let inputs, targets = batch_of_ids (load_ids !tokens) in
let original = Checkpoint.load !model in
let params = ref (Gpt2.of_checkpoint gpt2_124m original) in
let n_params =
let n = ref 0 in
Gpt2.Params.iter (fun _ -> incr n) !params;
!n
in
let initial_loss = repr_float (Nx.item [] (loss_fn inputs targets !params)) in
Printf.printf "initial loss (before step 1): %s\n%!" initial_loss;
let objective =
match !compute_dtype with
| "float32" -> loss_fn
| "bfloat16" -> loss_fn_half Nx.bfloat16
| "float16" -> loss_fn_half Nx.float16
| d ->
failwith
("--compute-dtype must be float32, bfloat16 or float16, got " ^ d)
in
let rate = !dropout in
if rate < 0.0 || rate >= 1.0 then failwith "--dropout must be in [0, 1)";
let obj key inputs targets params =
objective inputs targets
?dropout:(Option.map (fun k -> (rate, k)) key)
params
in
let root = Rune.Rng.key !seed in
let key_at i = if rate = 0.0 then None else Some (Rune.Rng.fold_in root i) in
let step : int -> Gpt2.t -> Gpt2.t * float =
if !devices = "" then
if !compute_dtype = "float16" then begin
let f =
Rune.jit2 ~device:!device
(module Scaled_in)
(module Scaled_out)
(train_step_scaled (fun key -> obj key inputs targets))
in
let ls = ref (Vega.Loss_scale.dynamic ()) in
fun i params ->
let out = f { Scaled_in.params; ls = !ls; key = key_at i } in
ls := out.Scaled_out.ls;
(out.Scaled_out.params, Nx.item [] out.Scaled_out.loss)
end
else
let f =
Rune.jit2 ~device:!device
(module Keyed_in)
(module Step_out)
(fun { Keyed_in.params; key } ->
train_step (obj key inputs targets) params)
in
fun i params ->
let out = f { Keyed_in.params; key = key_at i } in
(out.Step_out.params, Nx.item [] out.Step_out.loss)
else begin
if !compute_dtype = "float16" then
failwith "--compute-dtype float16 does not support --devices";
let devs = parse_devices !devices in
device := String.concat "," devs;
let in_axes =
List.init n_params (fun _ -> None)
@ [ Some 0; Some 0 ]
@ (if rate = 0.0 then [] else [ None ])
in
let f =
Rune.pmap2 ~devices:devs ~in_axes
(module Step_in)
(module Step_out)
(fun { Step_in.params; inputs; targets; key } ->
train_step (obj key inputs targets) params)
in
fun i params ->
let out = f { Step_in.params; inputs; targets; key = key_at i } in
(out.Step_out.params, Nx.item [] out.Step_out.loss)
end
in
let records = ref [] in
for i = 1 to !steps do
let t0 = Unix.gettimeofday () in
let params', loss = step i !params in
let ms = (Unix.gettimeofday () -. t0) *. 1e3 in
params := params';
Printf.printf "step %3d loss %s %9.2f ms\n%!" i (repr_float loss) ms;
let weights =
List.map (fun (n, t) -> (n, fingerprint t)) (fingerprint_tensors !params)
in
records := (i, repr_float loss, ms, weights) :: !records
done;
if !metrics_out <> "" then begin
let buf = Buffer.create (1 lsl 16) in
emit_metrics buf ~device:!device ~steps:!steps ~n_params ~initial_loss
(List.rev !records);
let oc = open_out !metrics_out in
Fun.protect
~finally:(fun () -> close_out oc)
(fun () -> output_string oc (Buffer.contents buf));
Printf.printf "wrote %s\n%!" !metrics_out
end;
if !save_weights <> "" then begin
Checkpoint.save !save_weights (checkpoint_of_params original !params);
Printf.printf "wrote %s\n%!" !save_weights
end