|
37 | 37 |
|
38 | 38 | CHECKPOINTS = modal.Volume.from_name("rababa-checkpoints") |
39 | 39 | 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") |
40 | 43 |
|
41 | 44 | 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 | + }, |
42 | 66 | "heb-diac-small": { |
43 | 67 | "teacher": "rababa_hebrew_byt5_s43/run-001/best", |
44 | 68 | "student_init": "google/byt5-small", |
@@ -292,7 +316,9 @@ def metrics(model) -> dict: |
292 | 316 |
|
293 | 317 | @app.local_entrypoint() |
294 | 318 | 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) |
296 | 322 | print(result) |
297 | 323 |
|
298 | 324 |
|
|
0 commit comments