Skip to content

Commit 3cbdcff

Browse files
committed
Supervise assistant wrapper spans in offset fallback
1 parent 9be126b commit 3cbdcff

2 files changed

Lines changed: 56 additions & 5 deletions

File tree

src/agentic_datagen/formatter.py

Lines changed: 53 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -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

403451
def _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

tests/test_formatter.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -476,7 +476,9 @@ def test_format_and_mask_uses_single_render_offset_mask_path_when_offsets_are_av
476476

477477
row = training_data[0]
478478
supervised_text = tokenizer.decode([token for token in row["labels"] if token != -100])
479-
assert supervised_text == "inspect repobashdone"
479+
assert supervised_text == "<assistant><think>inspect repo</think><tool_call>bash</tool_call></assistant><assistant>done</assistant>"
480+
assert "<tool_call>" in supervised_text
481+
assert "</think>" in supervised_text
480482
masked_text = tokenizer.decode(
481483
[token_id for token_id, label in zip(row["input_ids"], row["labels"]) if label == -100]
482484
)

0 commit comments

Comments
 (0)