6868 "hunk_id_per_token" ,
6969 "edit_op_per_token" ,
7070)
71+ _FAMILY_TOKEN_SIDE_CHANNEL_COLUMNS : Mapping [str , tuple [str , ...]] = {
72+ "semantic_graph" : _TOKEN_SEMANTIC_METADATA_COLUMNS ,
73+ "temporal_diff" : _TOKEN_TEMPORAL_METADATA_COLUMNS ,
74+ }
75+ _FAMILY_TOKEN_SIDE_CHANNEL_COLUMNS_FLAT = tuple (
76+ column
77+ for columns in _FAMILY_TOKEN_SIDE_CHANNEL_COLUMNS .values ()
78+ for column in columns
79+ )
7180_SOURCE_PLATFORM_IDS_COLUMN = "source_platform_ids"
7281_VALID_TOKEN_COUNT_COLUMN = "valid_token_count"
7382_ROW_METADATA_COLUMNS = (
@@ -179,6 +188,12 @@ class _SideChannelColumns:
179188 skipped : Sequence [Mapping [str , str | None ]]
180189
181190
191+ @dataclass (frozen = True )
192+ class _FamilySideChannelColumns :
193+ channels : Mapping [str , Mapping [str , np .ndarray ]]
194+ sources : Mapping [str , Mapping [str , Mapping [str , str | None ]]]
195+
196+
182197@dataclass (frozen = True )
183198class _ModelMetadataColumns :
184199 channels : Mapping [str , np .ndarray ]
@@ -288,6 +303,11 @@ def __init__(
288303 )
289304 side_channel_columns = _side_channel_windows (columns , token_rows , seq_len )
290305 side_channels = side_channel_columns .channels
306+ family_side_channel_columns = _family_side_channel_windows (
307+ columns ,
308+ token_rows ,
309+ seq_len ,
310+ )
291311 model_metadata_columns = _model_metadata_windows (columns , token_rows , seq_len )
292312 batch_metadata_columns = _resolve_batch_metadata_columns (
293313 columns ,
@@ -325,6 +345,13 @@ def __init__(
325345 key : _to_side_channel_values (key , value )
326346 for key , value in side_channels .items ()
327347 }
348+ self ._family_side_channels = {
349+ family : {
350+ column : value .astype (np .int32 , copy = False )
351+ for column , value in family_columns .items ()
352+ }
353+ for family , family_columns in family_side_channel_columns .channels .items ()
354+ }
328355 self ._model_metadata_channels = {
329356 key : value .astype (np .int32 , copy = False )
330357 for key , value in model_metadata_columns .channels .items ()
@@ -342,6 +369,7 @@ def __init__(
342369 document_id_source = document_id_source ,
343370 side_channel_sources = side_channel_columns .sources ,
344371 skipped_side_channels = side_channel_columns .skipped ,
372+ family_side_channel_sources = family_side_channel_columns .sources ,
345373 model_metadata_sources = model_metadata_columns .sources ,
346374 batch_metadata_columns = batch_metadata_columns ,
347375 )
@@ -441,6 +469,13 @@ def _make_batch(self, sample_idx: np.ndarray) -> LMTokenBatch:
441469 for key , value in self ._model_metadata_channels .items ()
442470 }
443471 )
472+ family_side_channels = {
473+ family : {
474+ column : mx .array (value [sample_idx ])
475+ for column , value in family_columns .items ()
476+ }
477+ for family , family_columns in self ._family_side_channels .items ()
478+ }
444479 return LMTokenBatch (
445480 tokens = mx .array (self ._tokens [sample_idx ]),
446481 target_tokens = None
@@ -452,6 +487,7 @@ def _make_batch(self, sample_idx: np.ndarray) -> LMTokenBatch:
452487 document_ids = None
453488 if self ._document_ids is None
454489 else mx .array (self ._document_ids [sample_idx ]),
490+ side_channels = family_side_channels or None ,
455491 metadata = self ._make_batch_metadata (sample_idx ),
456492 ** kwargs ,
457493 )
@@ -496,6 +532,7 @@ def _candidate_parquet_columns(
496532 candidates .append (text_key )
497533 candidates .extend (_LOSS_MASK_COLUMN_ALIASES )
498534 candidates .extend (_DOCUMENT_ID_COLUMN_ALIASES )
535+ candidates .extend (_FAMILY_TOKEN_SIDE_CHANNEL_COLUMNS_FLAT )
499536 candidates .append (_SOURCE_PLATFORM_IDS_COLUMN )
500537 candidates .append (_VALID_TOKEN_COUNT_COLUMN )
501538 for aliases in _SIDE_CHANNEL_COLUMN_ALIASES .values ():
@@ -721,6 +758,41 @@ def _side_channel_windows(
721758 return _SideChannelColumns (channels = channels , sources = sources , skipped = skipped )
722759
723760
761+ def _family_side_channel_windows (
762+ columns : ParquetColumns ,
763+ token_rows : list [list [int ]],
764+ seq_len : int ,
765+ ) -> _FamilySideChannelColumns :
766+ channels : dict [str , dict [str , np .ndarray ]] = {}
767+ sources : dict [str , dict [str , dict [str , str | None ]]] = {}
768+ for family , family_columns in _FAMILY_TOKEN_SIDE_CHANNEL_COLUMNS .items ():
769+ for column in family_columns :
770+ if column not in columns .values :
771+ continue
772+ _reject_non_integer_parquet_type (
773+ columns ,
774+ column ,
775+ f"{ column } { family } side-channel IDs" ,
776+ )
777+ rows = [
778+ _coerce_token_row (value , label = f"{ column } { family } side-channel" )
779+ for value in columns .require (column )
780+ ]
781+ if not _rows_are_token_aligned (rows , token_rows ):
782+ raise ValueError (
783+ f"{ column } rows must be token-aligned with token IDs"
784+ )
785+ channels .setdefault (family , {})[column ] = _fixed_windows_from_rows (
786+ rows ,
787+ seq_len ,
788+ )
789+ sources .setdefault (family , {})[column ] = {
790+ "column" : column ,
791+ "type" : columns .type_label (column ),
792+ }
793+ return _FamilySideChannelColumns (channels = channels , sources = sources )
794+
795+
724796def _model_metadata_windows (
725797 columns : ParquetColumns ,
726798 token_rows : list [list [int ]],
@@ -1137,6 +1209,10 @@ def _parquet_receipt(
11371209 document_id_source : str | None ,
11381210 side_channel_sources : Mapping [str , Mapping [str , str | None ]],
11391211 skipped_side_channels : Sequence [Mapping [str , str | None ]],
1212+ family_side_channel_sources : Mapping [
1213+ str ,
1214+ Mapping [str , Mapping [str , str | None ]],
1215+ ],
11401216 model_metadata_sources : Mapping [str , Mapping [str , str | None ]],
11411217 batch_metadata_columns : Sequence [str ],
11421218) -> dict [str , Any ]:
@@ -1173,6 +1249,14 @@ def _parquet_receipt(
11731249 receipt ["model_metadata_sources" ] = {
11741250 key : dict (value ) for key , value in sorted (model_metadata_sources .items ())
11751251 }
1252+ if family_side_channel_sources :
1253+ receipt ["family_side_channel_sources" ] = {
1254+ family : {
1255+ column : dict (source )
1256+ for column , source in sorted (columns .items ())
1257+ }
1258+ for family , columns in sorted (family_side_channel_sources .items ())
1259+ }
11761260 training_sources = {
11771261 "target_tokens" : target_source ,
11781262 "loss_mask" : loss_mask_source ,
0 commit comments