jvonrad commited on
Commit
07685c2
·
verified ·
1 Parent(s): 4642d81

Upload src/xscript/cli.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. src/xscript/cli.py +208 -0
src/xscript/cli.py ADDED
@@ -0,0 +1,208 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """`xscript` command-line entry point.
2
+
3
+ Pipeline order (see README):
4
+ flores-download -> byte-premium
5
+ tok-corpus -> tok-train -> tok-analyze (the tokenizer gate)
6
+ pool -> pack (per language x chosen tok)
7
+ train (one run of the matrix)
8
+ eval-bpb / eval-align -> bts (headline analysis)
9
+
10
+ Heavy steps are meant to run inside Slurm jobs (see slurm/); the CLI is the
11
+ single interface those jobs call, so behaviour is identical locally and on the
12
+ compute nodes.
13
+ """
14
+ import argparse
15
+
16
+ from .langs import (LANGS, TOK_FLAVORS, TOK_CONDITIONS, MODEL_FLAVORS,
17
+ tok_name, tok_conditions)
18
+
19
+
20
+ def _add(sub, name, help):
21
+ p = sub.add_parser(name, help=help)
22
+ return p
23
+
24
+
25
+ def main(argv=None):
26
+ ap = argparse.ArgumentParser(prog="xscript", description=__doc__,
27
+ formatter_class=argparse.RawDescriptionHelpFormatter)
28
+ sub = ap.add_subparsers(dest="cmd", required=True)
29
+
30
+ # ---- data prep ----
31
+ p = _add(sub, "flores-download", "download FLORES+ dev/devtest (needs HF_TOKEN)")
32
+ p.add_argument("--langs", nargs="*", default=list(LANGS))
33
+
34
+ _add(sub, "byte-premium", "compute FLORES+ byte premiums (+ compare Arnett)")
35
+
36
+ p = _add(sub, "tok-corpus", "build raw FineWeb/FineWeb2 tokenizer-training corpora")
37
+ p.add_argument("condition", choices=TOK_CONDITIONS + ["both"])
38
+ p.add_argument("--gb", type=float, default=4.0, help="target corpus size (GB)")
39
+
40
+ p = _add(sub, "tok-train", "train tokenizer(s): unigram/bpe/pa")
41
+ p.add_argument("--flavor", choices=TOK_FLAVORS + ["all"], default="all")
42
+ p.add_argument("--condition", choices=TOK_CONDITIONS + ["both"], default="both")
43
+
44
+ p = _add(sub, "tok-analyze", "fertility / allocation gate on FLORES+")
45
+ p.add_argument("--toks", nargs="*", default=None)
46
+
47
+ # ---- model-data prep ----
48
+ p = _add(sub, "pool", "build FineWeb(-2)-HQ text pool for a language")
49
+ p.add_argument("--lang", required=True, choices=list(LANGS))
50
+ p.add_argument("--gb", type=float, default=None, help="override byte budget (GB)")
51
+
52
+ p = _add(sub, "pack", "tokenize a pool into uint16 shards")
53
+ p.add_argument("--lang", required=True, choices=list(LANGS))
54
+ p.add_argument("--tok", required=True)
55
+ p.add_argument("--workers", type=int, default=8)
56
+
57
+ _add(sub, "plan", "print per-language pool budgets and the run matrix")
58
+
59
+ p = _add(sub, "runs", "list generated run names")
60
+ p.add_argument("--base", default="configs/base_main.yaml")
61
+ p.add_argument("--flavor", default="unigram", choices=MODEL_FLAVORS)
62
+ p.add_argument("--only-30b", action="store_true",
63
+ help="list 18 independent 30B runs (no extension trunks)")
64
+
65
+ # ---- training ----
66
+ p = _add(sub, "train", "train one run of the matrix")
67
+ p.add_argument("name", help="run name (see `xscript runs`)")
68
+ p.add_argument("--base", default="configs/base_main.yaml")
69
+ p.add_argument("--flavor", default="unigram", choices=MODEL_FLAVORS)
70
+ p.add_argument("--only-30b", action="store_true",
71
+ help="use a self-contained 30B WSD config, never a trunk branch")
72
+ p.add_argument("--output-name", default=None,
73
+ help="store an independent diagnostic replicate under this run name")
74
+ p.add_argument("--seed", type=int, default=None,
75
+ help="override model/optimizer RNG seed for a diagnostic replicate")
76
+ p.add_argument("--data-seed", type=int, default=None,
77
+ help="override packed-stream order seed for a diagnostic replicate")
78
+ p.add_argument("--wandb-id", default=None,
79
+ help="override the stable W&B run ID (useful for a clean replacement run)")
80
+
81
+ # ---- eval ----
82
+ p = _add(sub, "eval-bpb", "re-evaluate a checkpoint's BPB")
83
+ p.add_argument("name"); p.add_argument("--tok", required=True)
84
+ p.add_argument("--tag", default="final")
85
+
86
+ p = _add(sub, "eval-align", "MEXA alignment for a run")
87
+ p.add_argument("name"); p.add_argument("--tok", required=True)
88
+ p.add_argument("--split", default="dev")
89
+
90
+ p = _add(sub, "eval-bench", "downstream benchmarks (Global-MMLU/Belebele/XNLI) via lm-eval-harness")
91
+ p.add_argument("name"); p.add_argument("--tok", required=True)
92
+ p.add_argument("--tag", default="final")
93
+ p.add_argument("--tasks", nargs="*", default=None,
94
+ help="override tasks; default is all three benchmarks for the run's languages")
95
+ p.add_argument("--num-fewshot", type=int, default=0)
96
+ p.add_argument("--limit", type=float, default=None,
97
+ help="cap examples/task (for quick smoke checks)")
98
+ p.add_argument("--batch-size", type=int, default=4,
99
+ help="likelihood requests per GPU batch")
100
+ p.add_argument("--no-wandb", action="store_true")
101
+
102
+ p = _add(sub, "bts", "compute BTS + interaction across runs")
103
+ p.add_argument("--flavor", default="unigram", choices=MODEL_FLAVORS)
104
+ p.add_argument("--source", default="flores", choices=["flores", "holdout"])
105
+
106
+ args = ap.parse_args(argv)
107
+ return _dispatch(args)
108
+
109
+
110
+ def _dispatch(args):
111
+ cmd = args.cmd
112
+ if cmd == "flores-download":
113
+ from . import flores
114
+ flores.download(args.langs)
115
+
116
+ elif cmd == "byte-premium":
117
+ from . import byte_premium
118
+ byte_premium.run()
119
+
120
+ elif cmd == "tok-corpus":
121
+ from .data import tokcorpus
122
+ conds = TOK_CONDITIONS if args.condition == "both" else [args.condition]
123
+ for c in conds:
124
+ if c == "starved":
125
+ tokcorpus.build_starved(total_bytes=args.gb * 1e9)
126
+ else:
127
+ tokcorpus.build_destarved(total_bytes=args.gb * 1e9)
128
+
129
+ elif cmd == "tok-train":
130
+ from .tok import train as toktrain
131
+ flavors = TOK_FLAVORS if args.flavor == "all" else [args.flavor]
132
+ want = TOK_CONDITIONS if args.condition == "both" else [args.condition]
133
+ for f in flavors:
134
+ for c in want:
135
+ if c not in tok_conditions(f):
136
+ continue # pa has no starved condition
137
+ print(f"[tok-train] {tok_name(f, c)}")
138
+ toktrain.train(f, c)
139
+
140
+ elif cmd == "tok-analyze":
141
+ from .tok import analyze
142
+ analyze.run(args.toks)
143
+
144
+ elif cmd == "pool":
145
+ from .data import fineweb
146
+ budget = (args.gb * 1e9) if args.gb else fineweb.plan_budgets()[args.lang]
147
+ fineweb.build_pool(args.lang, budget)
148
+
149
+ elif cmd == "pack":
150
+ from .data import pack
151
+ pack.pack(args.lang, args.tok, workers=args.workers)
152
+
153
+ elif cmd == "plan":
154
+ _plan()
155
+
156
+ elif cmd == "runs":
157
+ from . import runmatrix
158
+ for n in runmatrix.list_runs(args.base, args.flavor, args.only_30b):
159
+ print(n)
160
+
161
+ elif cmd == "train":
162
+ from . import runmatrix, train
163
+ cfg = runmatrix.get_run(args.base, args.flavor, args.name, args.only_30b)
164
+ if args.output_name is not None:
165
+ cfg["name"] = args.output_name
166
+ if args.seed is not None:
167
+ cfg["seed"] = args.seed
168
+ if args.data_seed is not None:
169
+ cfg["data_seed"] = args.data_seed
170
+ if args.wandb_id is not None:
171
+ cfg["wandb_id"] = args.wandb_id
172
+ train.run_from_config(cfg)
173
+
174
+ elif cmd == "eval-bpb":
175
+ from .eval import bpb
176
+ bpb.run(args.name, args.tok, args.tag)
177
+
178
+ elif cmd == "eval-align":
179
+ from .eval import alignment
180
+ alignment.run(args.name, args.tok, args.split)
181
+
182
+ elif cmd == "eval-bench":
183
+ from .eval import bench
184
+ bench.run(args.name, args.tok, args.tag, tasks=args.tasks,
185
+ num_fewshot=args.num_fewshot, limit=args.limit,
186
+ log_wandb=not args.no_wandb, batch_size=args.batch_size)
187
+
188
+ elif cmd == "bts":
189
+ from .eval import bts
190
+ bts.run(args.flavor, args.source)
191
+
192
+
193
+ def _plan():
194
+ from .data.fineweb import plan_budgets
195
+ from . import runmatrix
196
+ b = plan_budgets()
197
+ print("Per-language pool byte budgets (worst-case, destarved tokenizer):")
198
+ for l, v in b.items():
199
+ print(f" {l}: {v/1e9:.1f} GB")
200
+ print("\nRun matrix (flavor=unigram):")
201
+ from . import _yaml
202
+ base = _yaml.load("configs/base_main.yaml")
203
+ for n in sorted(runmatrix.all_runs(base, "unigram")):
204
+ print(f" {n}")
205
+
206
+
207
+ if __name__ == "__main__":
208
+ main()