Skip to content

Commit 0231dd9

Browse files
author
Ronald Tse
committed
fix(distill): regenerate labels on torn resume-read instead of raising
The follow-on hole: when the Step-1 resume-read declares everything done, Step 2 re-reads the file — that read can itself serve a stale replica (6 valid pairs of 27,324 persisted), and raising there makes the watchdog relaunch into the identical state forever. On a torn read the run now regenerates all labels in-container and trains from the in-memory rows. Labeling extracted into a label_all() closure shared by the fresh/resume/regenerate paths.
1 parent 0286d5d commit 0231dd9

1 file changed

Lines changed: 64 additions & 52 deletions

File tree

‎src/gpu/modal_distill.py‎

Lines changed: 64 additions & 52 deletions
Original file line numberDiff line numberDiff line change
@@ -764,54 +764,52 @@ def collate(batch):
764764

765765
todo = [(s, t) for s, t in train_ds.rows if s not in done]
766766
fresh_rows: list[tuple[str, str]] = []
767-
if todo:
768-
print(f"[{spec_id}] labeling {len(todo)} remaining...", flush=True)
769-
770-
seq_max = int(spec.get("max_len", 384))
771-
772-
def label_batch(batch, max_len: int = 0):
773-
# lone-src OOM fallback truncates once, then skips: never
774-
# recurse on the same shape (torch 2.x renames the OOM
775-
# exception class, so match by message)
776-
if not max_len:
777-
max_len = seq_max
778-
try:
779-
enc = teacher_tok(
780-
[s for s, _ in batch],
781-
padding=True,
782-
truncation=True,
783-
max_length=max_len,
784-
return_tensors="pt",
785-
).to("cuda")
786-
with torch.inference_mode():
787-
out = teacher.generate(
788-
# r5 contract: generation cap = 2x window bytes
789-
# (diacritized output runs 1.4-1.6x input)
790-
**enc, max_new_tokens=2 * max_len, num_beams=label_beams
791-
)
792-
return [decode_joined(teacher_tok, o) for o in out]
793-
except RuntimeError as e:
794-
if "out of memory" not in str(e).lower():
795-
raise
796-
torch.cuda.empty_cache()
797-
if len(batch) == 1:
798-
if max_len > 128:
799-
return label_batch(batch, max_len=128)
800-
print(f" [{spec_id}] skipping pathological src", flush=True)
801-
return [None]
802-
mid = len(batch) // 2
803-
return label_batch(batch[:mid], max_len) + label_batch(
804-
batch[mid:], max_len
767+
seq_max = int(spec.get("max_len", 384))
768+
769+
def label_batch(batch, max_len: int = 0):
770+
# lone-src OOM fallback truncates once, then skips: never
771+
# recurse on the same shape (torch 2.x renames the OOM
772+
# exception class, so match by message)
773+
if not max_len:
774+
max_len = seq_max
775+
try:
776+
enc = teacher_tok(
777+
[s for s, _ in batch],
778+
padding=True,
779+
truncation=True,
780+
max_length=max_len,
781+
return_tensors="pt",
782+
).to("cuda")
783+
with torch.inference_mode():
784+
out = teacher.generate(
785+
# r5 contract: generation cap = 2x window bytes
786+
# (diacritized output runs 1.4-1.6x input)
787+
**enc, max_new_tokens=2 * max_len, num_beams=label_beams
805788
)
806-
789+
return [decode_joined(teacher_tok, o) for o in out]
790+
except RuntimeError as e:
791+
if "out of memory" not in str(e).lower():
792+
raise
793+
torch.cuda.empty_cache()
794+
if len(batch) == 1:
795+
if max_len > 128:
796+
return label_batch(batch, max_len=128)
797+
print(f" [{spec_id}] skipping pathological src", flush=True)
798+
return [None]
799+
mid = len(batch) // 2
800+
return label_batch(batch[:mid], max_len) + label_batch(
801+
batch[mid:], max_len
802+
)
803+
804+
def label_all(pairs: list[tuple[str, str]]) -> list[tuple[str, str]]:
807805
# deterministic token-budget batching: sort by length so long
808806
# srcs land in small batches — no OOM roulette
809-
todo.sort(key=lambda p: len(p[0].encode()))
807+
pairs = sorted(pairs, key=lambda p: len(p[0].encode()))
810808
budget = 32 * max(200, seq_max)
811809
batches: list[list[tuple[str, str]]] = []
812810
cur: list[tuple[str, str]] = []
813811
cur_max = 0
814-
for pair in todo:
812+
for pair in pairs:
815813
length = len(pair[0].encode())
816814
new_max = max(cur_max, length)
817815
if cur and (len(cur) + 1) * new_max > budget:
@@ -823,6 +821,7 @@ def label_batch(batch, max_len: int = 0):
823821
if cur:
824822
batches.append(cur)
825823

824+
rows: list[tuple[str, str]] = []
826825
labeled = 0
827826
with teacher_labels_path.open("a", encoding="utf-8") as fh:
828827
for batch in batches:
@@ -837,12 +836,12 @@ def label_batch(batch, max_len: int = 0):
837836
)
838837
+ "\n"
839838
)
840-
fresh_rows.append((src, text))
839+
rows.append((src, text))
841840
labeled += len(batch)
842841
if labeled <= 200 * 16 or labeled % 3200 < len(batch):
843842
mem = torch.cuda.memory_allocated() / 2**30
844843
print(
845-
f" labeled {labeled}/{len(todo)} (gpu {mem:.2f} GiB)",
844+
f" labeled {labeled}/{len(pairs)} (gpu {mem:.2f} GiB)",
846845
flush=True,
847846
)
848847
if labeled % 3200 < len(batch):
@@ -853,6 +852,11 @@ def label_batch(batch, max_len: int = 0):
853852
"rababa": CHECKPOINTS,
854853
"persian": PERSIAN_CHECKPOINTS,
855854
}.get(spec.get("out_volume", teacher_vol), SECRYST_CHECKPOINTS).commit()
855+
return rows
856+
857+
if todo:
858+
print(f"[{spec_id}] labeling {len(todo)} remaining...", flush=True)
859+
fresh_rows = label_all(todo)
856860
else:
857861
print(f"[{spec_id}] teacher labels already complete", flush=True)
858862

@@ -871,13 +875,7 @@ def accept_label(src: str, label: str) -> None:
871875
seen_labels.add(src)
872876
teacher_labels.append((src, label))
873877

874-
if fresh_rows:
875-
# this run generated the labels: use them directly. The volume
876-
# replica can serve a stale view of the just-written file (the
877-
# rababa/secrets tear: 2 visible pairs after 11,790 written).
878-
for src, label in fresh_rows:
879-
accept_label(src, label)
880-
else:
878+
if not fresh_rows:
881879
for line in teacher_labels_path.read_text(encoding="utf-8", errors="ignore").splitlines():
882880
if not line.strip():
883881
continue
@@ -886,11 +884,25 @@ def accept_label(src: str, label: str) -> None:
886884
except json.JSONDecodeError:
887885
continue # torn line from a volume replication race
888886
accept_label(row.get("src") or "", row.get("teacher") or "")
887+
if len(teacher_labels) < 0.5 * len(train_ds.rows):
888+
# stale replica of a complete file: regenerate rather than
889+
# fail (relaunch-only loops forever on this path)
890+
print(
891+
f"[{spec_id}] labels view torn ({len(teacher_labels)} valid "
892+
f"pairs); regenerating all labels",
893+
flush=True,
894+
)
895+
teacher_labels = []
896+
seen_labels = set()
897+
fresh_rows = label_all(list(train_ds.rows))
898+
if fresh_rows:
899+
for src, label in fresh_rows:
900+
accept_label(src, label)
889901
print(f"[{spec_id}] trainable label pairs: {len(teacher_labels)}", flush=True)
890902
if len(teacher_labels) < 0.5 * len(train_ds.rows):
891903
raise RuntimeError(
892-
f"labels view is torn: {len(teacher_labels)} valid pairs for "
893-
f"{len(train_ds.rows)} srcs — volume replication race; relaunch"
904+
f"labels view is torn even after regeneration: "
905+
f"{len(teacher_labels)} valid pairs for {len(train_ds.rows)} srcs"
894906
)
895907

896908
class TeacherPairs(Dataset):

0 commit comments

Comments
 (0)