Skip to content

Commit 9692ef3

Browse files
author
Ronald Tse
committed
feat(distill): Thai + Persian specs; sequence-level cross-tokenizer mode
Thai teacher is umt5 (sentencepiece) — logit KD is impossible across vocab spaces. Sequence-level mode: teacher generates labels with its own tokenizer (greedy, one pass), student trains CE on those labels with the byte tokenizer. Persian teacher is already byte-level, uses the existing logit-KD path. All four frozen teachers verified on their volumes.
1 parent fde9377 commit 9692ef3

1 file changed

Lines changed: 27 additions & 1 deletion

File tree

‎src/gpu/modal_distill.py‎

Lines changed: 27 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -37,8 +37,32 @@
3737

3838
CHECKPOINTS = modal.Volume.from_name("rababa-checkpoints")
3939
DATASETS = modal.Volume.from_name("rababa-datasets")
40+
SECRYST_CHECKPOINTS = modal.Volume.from_name("secryst-checkpoints")
41+
SECRYST_DATASETS = modal.Volume.from_name("secryst-datasets")
42+
PERSIAN_CHECKPOINTS = modal.Volume.from_name("persian-g2p-checkpoints")
4043

4144
SPECS: dict[str, dict[str, str]] = {
45+
"tha-g2p-small": {
46+
"teacher": "secryst_thai_ipa_thai_combined_mixed/run-001/best",
47+
"teacher_volume": "secryst",
48+
"student_init": "google/byt5-small",
49+
"train": "thai-ipa-expanded/train.jsonl",
50+
"val": "thai-ipa-expanded/val.jsonl",
51+
"test": "thai-ipa-expanded/test.jsonl",
52+
"out": "secryst_thai_g2p_distill_small/run-001",
53+
"mode": "sequence", # cross-tokenizer: teacher generates, student trains CE
54+
"note": "umt5 (sentencepiece) teacher -> ByT5-small byte student; +5pp PER gate",
55+
},
56+
"fas-g2p-small": {
57+
"teacher": "persian_g2p/run-001/best",
58+
"teacher_volume": "persian",
59+
"student_init": "google/byt5-small",
60+
"train": "persian_g2p/train.jsonl",
61+
"val": "persian_g2p/val.jsonl",
62+
"test": "persian_g2p/test.jsonl",
63+
"out": "interscript_fas_g2p_distill_small/run-001",
64+
"note": "ByT5-small teacher (already byte-level) -> ByT5-small student; CER gate",
65+
},
4266
"heb-diac-small": {
4367
"teacher": "rababa_hebrew_byt5_s43/run-001/best",
4468
"student_init": "google/byt5-small",
@@ -292,7 +316,9 @@ def metrics(model) -> dict:
292316

293317
@app.local_entrypoint()
294318
def main(spec: str = "heb-diac-small", epochs: int = 3) -> None:
295-
result = distill.remote(spec, epochs=epochs)
319+
mode = SPECS[spec].get("mode", "logit")
320+
fn = distill_sequence if mode == "sequence" else distill
321+
result = fn.remote(spec, epochs=epochs)
296322
print(result)
297323

298324

0 commit comments

Comments
 (0)