Skip to content

Commit 5db66c4

Browse files
author
Ronald Tse
committed
fix(distill): one spec->T5Config translation for all paths
The logit path's raw T5Config(**student_config) ignored the spec's enc_layers/dec_layers vocabulary (T5Config defaults won: a 6/6 student from a 6/4 spec) — the Hebrew layerdrop run crashed in layer_drop_state with IndexError and would otherwise have silently trained wrong depths. student_t5_config() is now shared by the sequence and logit paths.
1 parent e9cbe65 commit 5db66c4

2 files changed

Lines changed: 69 additions & 17 deletions

File tree

‎src/gpu/modal_distill.py‎

Lines changed: 25 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -170,6 +170,27 @@ def src_depth(prefix):
170170
out[name] = src.clone()
171171
return out
172172

173+
def student_t5_config(cfg: dict):
174+
"""Spec vocabulary -> T5Config, shared by every distill path.
175+
The spec speaks enc_layers/dec_layers and byte-model constants;
176+
T5Config speaks num_layers/num_decoder_layers. One translation."""
177+
from transformers import T5Config
178+
179+
return T5Config(
180+
vocab_size=259,
181+
d_model=cfg.get("d_model", 384),
182+
d_ff=cfg.get("d_ff", 1536),
183+
d_kv=cfg.get("d_kv", cfg.get("d_model", 384) // cfg.get("num_heads", 6)),
184+
num_layers=cfg.get("enc_layers", 8),
185+
num_decoder_layers=cfg.get("dec_layers", 8),
186+
num_heads=cfg.get("num_heads", 6),
187+
dropout_rate=0.1,
188+
feed_forward_proj=cfg.get("feed_forward_proj", "relu"),
189+
decoder_start_token_id=0,
190+
relative_attention_max_distance=128,
191+
)
192+
193+
173194
def _maybe_stitch(spec_id: str, spec: dict, student) -> None:
174195
"""When a custom-width student also names a pretrained init, bridge
175196
the pretrained weights down instead of random init (the capacity
@@ -307,9 +328,9 @@ def distill(spec_id: str, epochs: int = 3, alpha: float = 0.5, temperature: floa
307328
# tiny tier: no pretrained backbone at this width — random init
308329
# from an explicit config (dense teacher-label supervision, see
309330
# the spec note on collapse risk)
310-
from transformers import T5Config, T5ForConditionalGeneration
331+
from transformers import T5ForConditionalGeneration
311332

312-
cfg = T5Config(**spec["student_config"])
333+
cfg = student_t5_config(spec["student_config"])
313334
student = T5ForConditionalGeneration(cfg).to(device)
314335
if spec.get("layer_drop"):
315336
# depth-cut students keep the pretrained init: verbatim
@@ -719,22 +740,9 @@ def distill_sequence(spec_id: str, epochs: int = 3) -> dict:
719740
# teacher needs the GPU there; eviction-prone A10G headroom matters.
720741
student_tok = AutoTokenizer.from_pretrained("google/byt5-small")
721742
if spec.get("student_config"):
722-
from transformers import T5Config, T5ForConditionalGeneration
743+
from transformers import T5ForConditionalGeneration
723744

724-
cfg = spec["student_config"]
725-
config = T5Config(
726-
vocab_size=259,
727-
d_model=cfg.get("d_model", 384),
728-
d_ff=cfg.get("d_ff", 1536),
729-
d_kv=cfg.get("d_kv", cfg.get("d_model", 384) // cfg.get("num_heads", 6)),
730-
num_layers=cfg.get("enc_layers", 8),
731-
num_decoder_layers=cfg.get("dec_layers", 8),
732-
num_heads=cfg.get("num_heads", 6),
733-
dropout_rate=0.1,
734-
feed_forward_proj=cfg.get("feed_forward_proj", "relu"),
735-
decoder_start_token_id=0,
736-
relative_attention_max_distance=128,
737-
)
745+
config = student_t5_config(spec["student_config"])
738746
student = T5ForConditionalGeneration(config)
739747
_maybe_stitch(spec_id, spec, student)
740748
n_params = sum(q.numel() for q in student.parameters()) / 1e6

‎tests/test_student_config.py‎

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,44 @@
1+
"""student_t5_config: one translation of spec vocabulary -> T5Config,
2+
shared by the sequence and logit-KD paths (the logit path's raw
3+
T5Config(**spec) silently took T5 defaults for depths and produced a
4+
6/6 student from a 6/4 spec — the layer-copy then indexed out of
5+
range)."""
6+
7+
from __future__ import annotations
8+
9+
import sys
10+
from pathlib import Path
11+
12+
import pytest
13+
14+
sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "src"))
15+
16+
modal = pytest.importorskip("modal")
17+
18+
19+
def test_layerdrop_spec_translates_depths() -> None:
20+
from gpu.modal_distill import student_t5_config
21+
22+
cfg = student_t5_config(
23+
{
24+
"d_model": 1472,
25+
"d_kv": 64,
26+
"d_ff": 3584,
27+
"num_heads": 6,
28+
"enc_layers": 6,
29+
"dec_layers": 4,
30+
"feed_forward_proj": "gated-gelu",
31+
}
32+
)
33+
assert cfg.num_layers == 6
34+
assert cfg.num_decoder_layers == 4
35+
assert cfg.vocab_size == 259
36+
assert cfg.decoder_start_token_id == 0
37+
38+
39+
def test_byte_model_defaults() -> None:
40+
from gpu.modal_distill import student_t5_config
41+
42+
cfg = student_t5_config({})
43+
assert cfg.num_layers == 8 and cfg.num_decoder_layers == 8
44+
assert cfg.d_model == 384

0 commit comments

Comments
 (0)