Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
127 commits
Select commit Hold shift + click to select a range
b19565a
Reapply "Add MTP support for hybrid models (#2363)"
sancha Feb 2, 2026
f1401f1
Reapply "Fix two minor bugs in MTP implementation for hybrid models (…
sancha Feb 2, 2026
86eacb7
Handle one more missed reference
sancha Feb 2, 2026
1fa37a6
identity template and tool role mask
Oct 4, 2025
1bbdf4f
revert assert
Oct 8, 2025
3c0ca73
Skip empty sequences and chunks in MTP tensor roll
rkarimimahab Dec 24, 2025
8992c41
add option for SFTTokenizer to build_tokenizers
arendu Feb 4, 2026
cf3d0ea
resolve conflicts between args.sft and args.hybrid_context_parallel
arendu Feb 4, 2026
f3835d3
merged main
arendu Feb 5, 2026
8d5e58c
do not check for and add any EOD
arendu Feb 5, 2026
ec97314
Merge branch 'main' into adithyare/sft-ultra-v3-feb2026
arendu Feb 6, 2026
9ed15ae
Fix nan loss caused by zero token in MTP
BestJuly Feb 13, 2026
38dc023
Merge branch 'adithyare/sft-ultra-v3-feb2026' of https://github.com/a…
arendu Feb 13, 2026
fa016fd
rebased main
arendu Feb 18, 2026
e039525
identity template and tool role mask
Oct 4, 2025
ae4717d
revert assert
Oct 8, 2025
44c5c24
Skip empty sequences and chunks in MTP tensor roll
rkarimimahab Dec 24, 2025
8ed4905
add option for SFTTokenizer to build_tokenizers
arendu Feb 4, 2026
30f1f32
resolve conflicts between args.sft and args.hybrid_context_parallel
arendu Feb 4, 2026
72a2740
do not check for and add any EOD
arendu Feb 5, 2026
f82168a
removed dup identity template
arendu Feb 18, 2026
68514e8
added from #2363
arendu Feb 18, 2026
d10fa92
use core sft_tokenizer
arendu Feb 18, 2026
edd864c
fmt
arendu Feb 18, 2026
7bd7bfd
Merge remote-tracking branch 'github-upstream/main' into adithyare/sf…
arendu Feb 18, 2026
dc371e5
Merge branch adithyare/sft-ultra-v3-feb2026 from arendu/Megatron-LM
asolergi-nv Feb 19, 2026
7b44bc1
SFT: tokenizer, dataset, utils and sft_mamba updates (clean_thd)
asolergi-nv Feb 19, 2026
203fd61
add sft dataset to core datasets
asolergi-nv Feb 23, 2026
20c1841
Merge branch 'NVIDIA:main' into clean_thd
asolergi-nv Feb 23, 2026
38da580
create NullSFTTokenizer
asolergi-nv Feb 23, 2026
a4025c4
Add SFT dataset test
asolergi-nv Feb 23, 2026
cc54478
Merge branch 'NVIDIA:main' into clean_thd
asolergi-nv Feb 23, 2026
a4ae811
patch cu_seqlens
asolergi-nv Feb 23, 2026
a287001
add sft mamba script and refactor cp sharding logic
asolergi-nv Feb 23, 2026
4d47c63
Merge branch 'NVIDIA:main' into clean_thd
asolergi-nv Feb 24, 2026
f2ac856
remove logging - revert b4 merge
asolergi-nv Feb 24, 2026
bef1a64
create cu seqlens + remove some logging we should reintroduce b4 merge
asolergi-nv Feb 24, 2026
599a36d
small refactor to mtp_on_this_rank
asolergi-nv Feb 25, 2026
9211b08
Merge branch 'NVIDIA:main' into clean_thd
asolergi-nv Feb 25, 2026
af6385e
first step get batch on this tp rank
asolergi-nv Feb 25, 2026
e1669df
Merge origin/clean_thd (resolve sft_mamba conflict, keep local get_ba…
asolergi-nv Feb 25, 2026
3011ca9
Merge branch 'clean_thd' of https://github.com/asolergi-nv/Megatron-L…
asolergi-nv Feb 25, 2026
538bd5f
nit
asolergi-nv Feb 25, 2026
bdbf0f1
add pad token id for tests
asolergi-nv Feb 25, 2026
fc40a37
tests passing for sft & pretrain path
asolergi-nv Feb 25, 2026
58f8758
return batch elements alphabetically
asolergi-nv Feb 25, 2026
3de705b
Create position_ids in SFTDataset
asolergi-nv Feb 25, 2026
de79ab2
improve packing and refactor document sample shuffle indexes
asolergi-nv Mar 3, 2026
d039702
Merge branch 'main' into clean_thd
asolergi-nv Mar 3, 2026
1e72220
good packing and refactor test working
asolergi-nv Mar 3, 2026
f472ea8
renamed test
asolergi-nv Mar 3, 2026
e66741a
many fixes, now all tests passing with tp and cp
asolergi-nv Mar 3, 2026
cc626c7
add sft to preprocess data script
asolergi-nv Mar 5, 2026
c29cc9b
add completitions only training and update test
asolergi-nv Mar 5, 2026
ed5feb6
Merge branch 'NVIDIA:main' into clean_thd
asolergi-nv Mar 5, 2026
0d41a42
mock position ids fix
asolergi-nv Mar 5, 2026
ccb3704
cuseqlens in int32
asolergi-nv Mar 5, 2026
3075013
REVERT position ids and autotokenizer fix. Also fix for flash backend
asolergi-nv Mar 5, 2026
1a28a58
properly create position ids after padding & truncation like loss mask
asolergi-nv Mar 5, 2026
df867ab
add hybrid cp inputs to get batch method, update test, remove old cp …
asolergi-nv Mar 6, 2026
461bee8
Merge branch 'NVIDIA:main' into clean_thd
asolergi-nv Mar 9, 2026
a19d6a7
commiting cp stuff but reverting
asolergi-nv Mar 9, 2026
691a755
delete attention mask from hybridcp path and note in hybird cp test t…
asolergi-nv Mar 9, 2026
a891e8c
some cleaning
asolergi-nv Mar 9, 2026
d80b55a
updated test. more checks and vibecoded hybrid cp tests
asolergi-nv Mar 9, 2026
4974826
add docstring to utils funcs
asolergi-nv Mar 11, 2026
8641b99
add sft data inspector script
asolergi-nv Mar 12, 2026
80da7b2
Merge origin/main into clean_thd, accept main's gpt_model.py MTP changes
asolergi-nv Mar 12, 2026
9ee5dc2
add tool support, add loss mask with proper masking, improve get batc…
asolergi-nv Mar 12, 2026
e664a03
Merge branch 'NVIDIA:main' into clean_thd
asolergi-nv Mar 13, 2026
200d3e5
improve SFTDataset tests
asolergi-nv Mar 13, 2026
69a9e0a
Merge branch 'clean_thd' of https://github.com/asolergi-nv/Megatron-L…
asolergi-nv Mar 16, 2026
33bc799
Merge upstream/main into clean_thd
asolergi-nv Mar 16, 2026
f00322f
Merge remote-tracking branch 'upstream/main' into clean_thd
asolergi-nv Mar 17, 2026
ffdbb99
Merge branch 'NVIDIA:main' into clean_thd
asolergi-nv Mar 18, 2026
b0a76ba
fix conflicts
asolergi-nv Mar 18, 2026
b51203e
Merge branch 'NVIDIA:main' into clean_thd
asolergi-nv Mar 18, 2026
363b288
fix missing tool_response_end_tokens and remove adding special tokens…
asolergi-nv Mar 18, 2026
6dfbc23
remove megatron training SFTDataset and move IGNORE_INDEX
asolergi-nv Mar 18, 2026
ee02362
redo build tokenizer
asolergi-nv Mar 18, 2026
c523331
fix think tokens tokenize and add knob for add_special_tokens
asolergi-nv Mar 19, 2026
8fdb933
Merge branch 'NVIDIA:main' into clean_thd
asolergi-nv Mar 19, 2026
949cf73
cleaning
asolergi-nv Mar 19, 2026
ae7f1f2
autoformat linting
asolergi-nv Mar 19, 2026
4999b40
update pretrain mamba and gpt entry points
asolergi-nv Mar 19, 2026
4ad9e80
Merge branch 'NVIDIA:main' into clean_thd
asolergi-nv Mar 19, 2026
9f0e3a6
add copyright headers
asolergi-nv Mar 19, 2026
e449719
fix ci install
asolergi-nv Mar 19, 2026
01e0a5f
fix ci
asolergi-nv Mar 19, 2026
01481d3
Merge branch 'main' into clean_thd
asolergi-nv Mar 19, 2026
f1a596e
More CI fixes
asolergi-nv Mar 19, 2026
c832039
Merge branch 'main' into clean_thd
asolergi-nv Mar 19, 2026
12ec5e6
lint
asolergi-nv Mar 19, 2026
c766445
remove null sft tokenizer
asolergi-nv Mar 19, 2026
f86ea8a
fix more ci errors
asolergi-nv Mar 20, 2026
45b402d
test
asolergi-nv Mar 20, 2026
be6f0c8
Merge branch 'main' into clean_thd
asolergi-nv Mar 20, 2026
4dddd94
lint
asolergi-nv Mar 20, 2026
a61dd23
ci
asolergi-nv Mar 20, 2026
7740a2c
Merge branch 'main' into clean_thd
asolergi-nv Mar 20, 2026
cb5e4fb
fix hybridcp cp sharding
asolergi-nv Mar 20, 2026
5544ee2
Merge branch 'main' into clean_thd
asolergi-nv Mar 20, 2026
e785d9b
What a bug
asolergi-nv Mar 20, 2026
b9fb498
nits
asolergi-nv Mar 20, 2026
6302c75
Merge branch 'main' into clean_thd
asolergi-nv Mar 23, 2026
f80b452
Merge branch 'main' into clean_thd
asolergi-nv Mar 24, 2026
3564e8d
nits before review
asolergi-nv Mar 24, 2026
c864503
doc nit and torch seeding
asolergi-nv Mar 24, 2026
4c87a7c
work with numpy arrays directly
asolergi-nv Mar 24, 2026
74c9b6c
delete get_batch_on_this_tp_rank from megatron training utils
asolergi-nv Mar 24, 2026
39d1623
force cp > 1 for hybrid cp tests
asolergi-nv Mar 24, 2026
24adf0d
Apply suggestions from code review
asolergi-nv Mar 24, 2026
33d9372
Merge branch 'main' into clean_thd
asolergi-nv Mar 24, 2026
e969aa1
lint
asolergi-nv Mar 24, 2026
ac0cf2a
nits
asolergi-nv Mar 24, 2026
3e1ddf7
Merge branch 'main' into clean_thd
asolergi-nv Mar 24, 2026
754aa2d
lint
asolergi-nv Mar 24, 2026
957911b
fix cp calls
asolergi-nv Mar 24, 2026
d6c9b29
Merge branch 'main' into clean_thd
asolergi-nv Mar 24, 2026
e450e34
Merge branch 'main' into clean_thd
asolergi-nv Mar 25, 2026
f07ae4b
nit
asolergi-nv Mar 25, 2026
b6e48b1
Merge branch 'main' into clean_thd
asolergi-nv Mar 29, 2026
dd16ebe
nit
asolergi-nv Mar 30, 2026
c093244
Merge branch 'main' into clean_thd
asolergi-nv Mar 30, 2026
2b8e74f
more hope
asolergi-nv Mar 30, 2026
807ee2e
Merge branch 'main' into clean_thd
asolergi-nv Mar 30, 2026
8b7126b
Merge branch 'main' into clean_thd
asolergi-nv Apr 1, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion examples/multimodal/train.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,8 @@
is_pipeline_last_stage,
)
from megatron.training import get_args, get_timers, get_tokenizer, pretrain
from megatron.training.utils import is_last_rank, get_batch_on_this_cp_rank
from megatron.core.utils import get_batch_on_this_cp_rank
from megatron.training.utils import is_last_rank


def get_batch(data_iterator, image_token_index, img_seq_len):
Expand Down
4 changes: 3 additions & 1 deletion examples/post_training/modelopt/finetune.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,8 @@
)
from utils import get_hf_tokenizer
from model_provider import model_provider
from megatron.core.parallel_state import get_context_parallel_group


REMOVE_THINK_CHAT_TEMPLATE = (
"{% if '</think>' in content %}{% set content = content.split('</think>')[-1] %}{% endif %}"
Expand Down Expand Up @@ -435,7 +437,7 @@ def get_batch(data_iterator):
batch["hidden_states"] = feature_b["hidden_states"].transpose(0, 1)[:args.seq_length]

# slice batch along sequence dimension for context parallelism
batch = get_batch_on_this_cp_rank(batch)
batch = get_batch_on_this_cp_rank(batch, is_hybrid_cp=False, cp_group=get_context_parallel_group())

return batch

Expand Down
5 changes: 4 additions & 1 deletion examples/post_training/modelopt/quantize.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,7 @@
mtq_luts = None
warnings.warn("luts is not installed. LUTs quantization configs will not be available.")

from megatron.core.parallel_state import get_context_parallel_group
from megatron.core.utils import get_batch_on_this_cp_rank
from megatron.post_training.arguments import add_modelopt_args
from megatron.post_training.checkpointing import load_modelopt_checkpoint
Expand Down Expand Up @@ -410,7 +411,9 @@ def _dataset_forward_loop_func(model):
batch_size=args.calib_batch_size,
)
for sample in tqdm(dataloader, disable=torch.distributed.get_rank()):
sample = get_batch_on_this_cp_rank(sample)
sample = get_batch_on_this_cp_rank(
sample, is_hybrid_cp=False, cp_group=get_context_parallel_group()
)
simple_generate(model, sample["input_ids"], osl=1, calibration_mode=True)

unwrapped_model = unwrap_model(model)[0]
Expand Down
332 changes: 332 additions & 0 deletions inspect_sft_file_prefix.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,332 @@
# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.

#!/usr/bin/env python3
"""Inspect pretokenized SFT samples from Megatron-LM .bin/.idx file pairs.

Usage:
python inspect_sft_file_prefix.py <file_prefix>

Example:
python inspect_sft_file_prefix.py /path/to/dataset-materialized_text_document
"""

import argparse
import struct
import numpy
from functools import lru_cache
from typing import Optional, Tuple

# ──────────────────────────────────────────────────────────────────────────────
# Hardcoded config for nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-FP8 chat template
# ──────────────────────────────────────────────────────────────────────────────
TOKENIZER_NAME = "nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-FP8"

# These are pre-computed from:
# tokenizer.encode("<|im_start|>system\n", add_special_tokens=False) etc.
ROLE_START_TOKENS = {
"system": [10, 25708, 1010], # <|im_start|>system\n
"user": [10, 3263, 1010], # <|im_start|>user\n
"assistant": [10, 1503, 19464, 1010], # <|im_start|>assistant\n
}
END_TOKENS = [11, 1010] # <|im_end|>\n
THINK_START_ID = [12] # <think>
THINK_END_ID = [13] # </think>
TOOL_CALL_START = [14, 1010] # <tool_call>\n
TOOL_CALL_END = [15, 1010] # </tool_call>\n
TOOL_RESPONSE_START = [16, 1010] # <tool_response>\n
TOOL_RESPONSE_END = [17, 1010] # </tool_response>\n

# ──────────────────────────────────────────────────────────────────────────────
# Index / bin readers (from Megatron-LM)
# ──────────────────────────────────────────────────────────────────────────────
_INDEX_HEADER = b"MMIDIDX\x00\x00"


class _MMapBinReader:
def __init__(self, bin_path: str) -> None:
self._bin_file_reader = open(bin_path, mode="rb")
self._bin_buffer_mmap = numpy.memmap(self._bin_file_reader, mode="r", order="C")
self._bin_buffer = memoryview(self._bin_buffer_mmap.data)

def read(self, dtype, count: int, offset: int) -> numpy.ndarray:
return numpy.frombuffer(self._bin_buffer, dtype=dtype, count=count, offset=offset)

def __del__(self) -> None:
if self._bin_buffer_mmap is not None:
self._bin_buffer_mmap._mmap.close()
if self._bin_file_reader is not None:
self._bin_file_reader.close()
del self._bin_buffer_mmap
del self._bin_file_reader


class _IndexReader:
def __init__(self, idx_path: str) -> None:
with open(idx_path, "rb") as stream:
header = stream.read(9)
assert header == _INDEX_HEADER, f"bad header, cannot read: {idx_path}"

version = struct.unpack("<Q", stream.read(8))[0]
assert version == 1, f"bad version, cannot read: {idx_path}"

_code = struct.unpack("<B", stream.read(1))[0]
self.sequence_count = struct.unpack("<Q", stream.read(8))[0]
self.document_count = struct.unpack("<Q", stream.read(8))[0]
offset = stream.tell()

self.bin_buffer_mmap = numpy.memmap(idx_path, mode="r", order="C")
self.bin_buffer = memoryview(self.bin_buffer_mmap)

self.sequence_lengths = numpy.frombuffer(
self.bin_buffer, dtype=numpy.int32, count=self.sequence_count, offset=offset
)
self.sequence_pointers = numpy.frombuffer(
self.bin_buffer,
dtype=numpy.int64,
count=self.sequence_count,
offset=offset + self.sequence_lengths.nbytes,
)
self.document_indices = numpy.frombuffer(
self.bin_buffer,
dtype=numpy.int64,
count=self.document_count,
offset=offset + self.sequence_lengths.nbytes + self.sequence_pointers.nbytes,
)

def __len__(self) -> int:
return self.sequence_count

@lru_cache(maxsize=8)
def __getitem__(self, idx: int) -> Tuple[numpy.int64, numpy.int32]:
return (self.sequence_pointers[idx], self.sequence_lengths[idx])

def __del__(self) -> None:
if hasattr(self, "bin_buffer_mmap"):
self.bin_buffer_mmap._mmap.close()
del self.bin_buffer_mmap


# ──────────────────────────────────────────────────────────────────────────────
# Segment extraction logic
# ──────────────────────────────────────────────────────────────────────────────
def find_subsequence(sequence, subsequence, start=0):
sub_len = len(subsequence)
for i in range(start, len(sequence) - sub_len + 1):
if sequence[i : i + sub_len] == subsequence:
return i
return -1


NL_TOKEN = 1010 # \n token id

def split_tool_calls(tokens, offset):
"""Split a token sequence into assistant text and tool_call sub-segments.

Whitespace-only assistant fragments (e.g. a lone \\n between </think> and
<tool_call>) are dropped so we don't produce meaningless segments.
"""
tc_start_len = len(TOOL_CALL_START)
tc_end_len = len(TOOL_CALL_END)
result = []
pos = 0
while pos < len(tokens):
# Find next <tool_call>\n
tc_start = find_subsequence(tokens, TOOL_CALL_START, pos)

if tc_start == -1:
# No more tool calls, rest is regular assistant content
if pos < len(tokens):
result.append({"role": "assistant", "tokens": tokens[pos:], "start": offset + pos, "end": offset + len(tokens)})
break

# Assistant content before tool_call (skip if whitespace-only)
if tc_start > pos:
frag = tokens[pos:tc_start]
if not all(t == NL_TOKEN for t in frag):
result.append({"role": "assistant", "tokens": frag, "start": offset + pos, "end": offset + tc_start})

# Find matching </tool_call>\n
content_start = tc_start + tc_start_len
tc_end = find_subsequence(tokens, TOOL_CALL_END, content_start)

if tc_end == -1:
# No closing tag, treat rest as tool_call
result.append({"role": "tool_call", "tokens": tokens[content_start:], "start": offset + content_start, "end": offset + len(tokens)})
break

# Tool call content (excluding markers)
result.append({"role": "tool_call", "tokens": tokens[content_start:tc_end], "start": offset + content_start, "end": offset + tc_end})
pos = tc_end + tc_end_len

# Also drop trailing whitespace-only assistant fragments
if result and result[-1]["role"] == "assistant" and all(t == NL_TOKEN for t in result[-1]["tokens"]):
result.pop()

return result


def extract_segments(tokenized_conversation, role_start_tokens, end_tokens, think_start_id, think_end_id):
markers = []
for role, start_tokens in role_start_tokens.items():
pos = 0
while True:
idx = find_subsequence(tokenized_conversation, start_tokens, pos)
if idx == -1:
break
markers.append((idx, role, len(start_tokens)))
pos = idx + len(start_tokens)
markers.sort(key=lambda x: x[0])

segments = []
for start_pos, role, marker_len in markers:
content_start = start_pos + marker_len
end_pos = find_subsequence(tokenized_conversation, end_tokens, content_start)
if end_pos == -1:
content_end = len(tokenized_conversation)
else:
content_end = end_pos
content_tokens = tokenized_conversation[content_start:content_end]

# Check if this user turn is actually a tool response
if role == "user" and len(content_tokens) >= len(TOOL_RESPONSE_START) and content_tokens[:len(TOOL_RESPONSE_START)] == TOOL_RESPONSE_START:
segments.append({"role": "tool_response", "tokens": content_tokens, "start": content_start, "end": content_end})
continue

if role == "assistant":
think_start_idx = find_subsequence(content_tokens, think_start_id)
if think_start_idx != -1:
think_end_idx = find_subsequence(content_tokens, think_end_id, think_start_idx + len(think_start_id))
else:
think_end_idx = -1

if think_start_idx != -1 and think_end_idx != -1:
reasoning_tokens = content_tokens[think_start_idx + len(think_start_id) : think_end_idx]
response_tokens = content_tokens[think_end_idx + len(think_end_id) :]
if reasoning_tokens:
abs_start = content_start + think_start_idx + len(think_start_id)
abs_end = content_start + think_end_idx
segments.append({"role": "reasoning", "tokens": reasoning_tokens, "start": abs_start, "end": abs_end})
# Split the response part by tool calls
if response_tokens:
abs_start = content_start + think_end_idx + len(think_end_id)
segments.extend(split_tool_calls(response_tokens, abs_start))
continue

# No think tags — split entire content by tool calls
if find_subsequence(content_tokens, TOOL_CALL_START) != -1:
segments.extend(split_tool_calls(content_tokens, content_start))
continue

segments.append({"role": role, "tokens": content_tokens, "start": content_start, "end": content_end})

return segments


# ──────────────────────────────────────────────────────────────────────────────
# Main
# ──────────────────────────────────────────────────────────────────────────────
def main():
parser = argparse.ArgumentParser(description="Inspect pretokenized SFT samples.")
parser.add_argument("file_prefix", help="Path prefix for .bin/.idx files (without extension)")
args = parser.parse_args()

print(f"Loading tokenizer: {TOKENIZER_NAME}")
from transformers import AutoTokenizer
tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_NAME, trust_remote_code=True)

# Verify hardcoded token ids match this tokenizer
assert tokenizer.encode("<|im_start|>system\n", add_special_tokens=False) == ROLE_START_TOKENS["system"]
assert tokenizer.encode("<|im_start|>user\n", add_special_tokens=False) == ROLE_START_TOKENS["user"]
assert tokenizer.encode("<|im_start|>assistant\n", add_special_tokens=False) == ROLE_START_TOKENS["assistant"]
assert tokenizer.encode("<|im_end|>\n", add_special_tokens=False) == END_TOKENS
assert tokenizer.encode("<think>", add_special_tokens=False) == THINK_START_ID
assert tokenizer.encode("</think>", add_special_tokens=False) == THINK_END_ID
assert tokenizer.encode("<tool_call>\n", add_special_tokens=False) == TOOL_CALL_START
assert tokenizer.encode("</tool_call>\n", add_special_tokens=False) == TOOL_CALL_END
assert tokenizer.encode("<tool_response>\n", add_special_tokens=False) == TOOL_RESPONSE_START
assert tokenizer.encode("</tool_response>\n", add_special_tokens=False) == TOOL_RESPONSE_END

print(f"Loading index: {args.file_prefix}.idx")
index = _IndexReader(args.file_prefix + ".idx")
print(f"Loading bin: {args.file_prefix}.bin")
reader = _MMapBinReader(args.file_prefix + ".bin")
print(f"Total sequences: {len(index)}\n")

while True:
try:
raw = input(f"Enter sample index [0-{len(index) - 1}] (q to quit): ").strip()
except (EOFError, KeyboardInterrupt):
print()
break

if raw.lower() == "q":
break

# Parse optional -d suffix for raw detokenized output
detokenize_only = raw.endswith("-d")
if detokenize_only:
raw = raw[:-2].strip()

try:
sample_idx = int(raw)
except ValueError:
print(f"Invalid input: {raw!r}")
continue

if sample_idx < 0 or sample_idx >= len(index):
print(f"Out of range. Must be 0-{len(index) - 1}")
continue

pointer, length = index[sample_idx]
sample = reader.read(numpy.int32, int(length), int(pointer))

if detokenize_only:
print(f"\n{'=' * 80}")
print(f"Sample {sample_idx} | {len(sample)} total tokens (raw detokenized)")
print(f"{'=' * 80}\n")
print(tokenizer.decode(sample.tolist()))
print(f"\n{'=' * 80}\n")
continue

segments = extract_segments(
sample.tolist(), ROLE_START_TOKENS, END_TOKENS, THINK_START_ID, THINK_END_ID
)

print(f"\n{'=' * 80}")
print(f"Sample {sample_idx} | {len(sample)} total tokens")
print(f"{'=' * 80}")

assistant_tokens = 0
reasoning_tokens = 0
tool_call_tokens = 0

for seg in segments:
decoded = tokenizer.decode(seg["tokens"])
role_label = seg["role"].upper()
n_tokens = len(seg["tokens"])
print(f"\n[{role_label:>13}] ({n_tokens:>5} tokens | {seg['start']}:{seg['end']})")
print(decoded)

if seg["role"] == "assistant":
assistant_tokens += n_tokens
elif seg["role"] == "reasoning":
reasoning_tokens += n_tokens
elif seg["role"] == "tool_call":
tool_call_tokens += n_tokens

has_tools = tool_call_tokens > 0
has_reasoning = reasoning_tokens > 0

print(f"\n{'-' * 80}")
print(f"Training tokens (assistant only): {assistant_tokens}")
if has_tools:
print(f"Training tokens (assistant + tool calls): {assistant_tokens + tool_call_tokens}")
if has_reasoning:
print(f"Training tokens (assistant + reasoning): {assistant_tokens + reasoning_tokens}")
if has_tools or has_reasoning:
print(f"Training tokens (assistant + tool calls + reasoning): {assistant_tokens + tool_call_tokens + reasoning_tokens}")
print(f"{'=' * 80}\n")


if __name__ == "__main__":
main()
Loading
Loading