@@ -389,15 +389,63 @@ def _strip_markers_and_collect_spans(text: str, markers: list[tuple[str, str]])
389389 cleaned_text = "" .join (cleaned_parts )
390390 if not spans :
391391 return cleaned_text , []
392- spans .sort ()
393- merged_spans : list [tuple [int , int ]] = [spans [0 ]]
394- for start , end in spans [1 :]:
392+ return cleaned_text , _merge_spans (spans )
393+
394+
395+ def _merge_spans (spans : list [tuple [int , int ]]) -> list [tuple [int , int ]]:
396+ if not spans :
397+ return []
398+ ordered_spans = sorted (spans )
399+ merged_spans : list [tuple [int , int ]] = [ordered_spans [0 ]]
400+ for start , end in ordered_spans [1 :]:
395401 last_start , last_end = merged_spans [- 1 ]
396402 if start <= last_end :
397403 merged_spans [- 1 ] = (last_start , max (last_end , end ))
398404 else :
399405 merged_spans .append ((start , end ))
400- return cleaned_text , merged_spans
406+ return merged_spans
407+
408+
409+ def _expand_span_to_assistant_block (text : str , start : int , end : int ) -> tuple [int , int ] | None :
410+ assistant_start_tokens = (
411+ "<|im_start|>assistant\n " ,
412+ "<|assistant|>\n " ,
413+ "<|assistant|>" ,
414+ "<assistant>" ,
415+ )
416+ assistant_end_tokens = (
417+ "<|im_end|>" ,
418+ "</assistant>" ,
419+ "</s>" ,
420+ )
421+ block_start = - 1
422+ for token in assistant_start_tokens :
423+ token_start = text .rfind (token , 0 , start )
424+ if token_start > block_start :
425+ block_start = token_start
426+ if block_start < 0 :
427+ return None
428+ block_end = - 1
429+ for token in assistant_end_tokens :
430+ token_end_start = text .find (token , end )
431+ if token_end_start >= 0 and (block_end < 0 or token_end_start < block_end ):
432+ block_end = token_end_start + len (token )
433+ if block_end < 0 :
434+ return None
435+ while block_end < len (text ) and text [block_end ] in "\r \n " :
436+ block_end += 1
437+ return block_start , block_end
438+
439+
440+ def _expand_supervised_spans (text : str , supervised_spans : list [tuple [int , int ]]) -> list [tuple [int , int ]]:
441+ expanded_spans : list [tuple [int , int ]] = []
442+ for start , end in supervised_spans :
443+ expanded_span = _expand_span_to_assistant_block (text , start , end )
444+ if expanded_span is None :
445+ expanded_spans .append ((start , end ))
446+ else :
447+ expanded_spans .append (expanded_span )
448+ return _merge_spans (expanded_spans )
401449
402450
403451def _labels_from_offsets (
@@ -438,6 +486,7 @@ def _offset_mask_row(
438486 if stripped is None :
439487 return None
440488 formatted_text , supervised_spans = stripped
489+ supervised_spans = _expand_supervised_spans (formatted_text , supervised_spans )
441490 encoded = _tokenize_text_with_offsets (text_tokenizer , formatted_text )
442491 if encoded is None :
443492 return None
0 commit comments