@@ -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