Skip to content

Commit 8758a3e

Browse files
author
Ronald Tse
committed
fix(distill): label length cap follows spec max_len; teacher stays on GPU until labels are usable
The Step-2 label filter hard-rejected anything over 384 bytes — with Arabic's 1450-byte windows that discarded essentially every label (8 of 11,793 survived). This, not volume replication, was the original '2-6 valid pairs' disease. Cap is now 2x spec max_len. Regeneration also ran after the teacher had been moved to CPU — the confirmation step now precedes the free.
1 parent b94f399 commit 8758a3e

1 file changed

Lines changed: 14 additions & 11 deletions

File tree

‎src/gpu/modal_distill.py‎

Lines changed: 14 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -898,18 +898,16 @@ def label_all(pairs: list[tuple[str, str]]) -> list[tuple[str, str]]:
898898
else:
899899
print(f"[{spec_id}] teacher labels already complete", flush=True)
900900

901-
# Step 2: student trains on teacher labels (teacher no longer
902-
# needed on GPU — free it before the training loop)
903-
teacher.to("cpu")
904-
torch.cuda.empty_cache()
905-
student.to("cuda")
906-
student.gradient_checkpointing_enable()
901+
# Step 2: student trains on teacher labels. The teacher stays on the
902+
# GPU until the labels are confirmed usable — regeneration (below)
903+
# still needs it.
904+
label_cap = 2 * int(spec.get("max_len", 384))
907905
teacher_labels = []
908906
seen_labels: set[str] = set()
909907

910908
def accept_label(src: str, label: str) -> None:
911909
src, label = src.strip(), label.strip()
912-
if src and src not in seen_labels and label and len(label.encode()) <= 384:
910+
if src and src not in seen_labels and label and len(label.encode()) <= label_cap:
913911
seen_labels.add(src)
914912
teacher_labels.append((src, label))
915913

@@ -923,10 +921,10 @@ def accept_label(src: str, label: str) -> None:
923921
continue # torn line from a volume replication race
924922
accept_label(row.get("src") or "", row.get("teacher") or "")
925923
if len(teacher_labels) < 0.5 * len(train_ds.rows):
926-
# stale replica of a complete file: regenerate rather than
927-
# fail (relaunch-only loops forever on this path)
924+
# the file is unusable from this container: regenerate rather
925+
# than fail (relaunch-only loops forever on this path)
928926
print(
929-
f"[{spec_id}] labels view torn ({len(teacher_labels)} valid "
927+
f"[{spec_id}] labels unusable ({len(teacher_labels)} valid "
930928
f"pairs); regenerating all labels",
931929
flush=True,
932930
)
@@ -939,10 +937,15 @@ def accept_label(src: str, label: str) -> None:
939937
print(f"[{spec_id}] trainable label pairs: {len(teacher_labels)}", flush=True)
940938
if len(teacher_labels) < 0.5 * len(train_ds.rows):
941939
raise RuntimeError(
942-
f"labels view is torn even after regeneration: "
940+
f"labels unusable even after regeneration: "
943941
f"{len(teacher_labels)} valid pairs for {len(train_ds.rows)} srcs"
944942
)
945943

944+
teacher.to("cpu")
945+
torch.cuda.empty_cache()
946+
student.to("cuda")
947+
student.gradient_checkpointing_enable()
948+
946949
class TeacherPairs(Dataset):
947950
def __len__(self):
948951
return len(teacher_labels)

0 commit comments

Comments
 (0)