Skip to content

Commit 3fa84fc

Browse files
committed
merge: update branch to latest main
Signed-off-by: yaoyu-33 <yaoyu.094@gmail.com>
2 parents 0e9a909 + 18ae6f5 commit 3fa84fc

5 files changed

Lines changed: 60 additions & 4 deletions

File tree

docs/versions1.json

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -4,14 +4,15 @@
44
"url": "https://docs.nvidia.com/nemo/megatron-bridge/nightly/"
55
},
66
{
7+
"name": "0.6.0 (latest) · 26.08",
78
"version": "0.6.0",
8-
"url": "https://docs.nvidia.com/nemo/megatron-bridge/0.6.0/"
9+
"url": "https://docs.nvidia.com/nemo/megatron-bridge/0.6.0/",
10+
"preferred": true
911
},
1012
{
1113
"name": "0.5.1 (latest) · 26.06.01",
1214
"version": "0.5.1",
13-
"url": "https://docs.nvidia.com/nemo/megatron-bridge/0.5.1/",
14-
"preferred": true
15+
"url": "https://docs.nvidia.com/nemo/megatron-bridge/0.5.1/"
1516
},
1617
{
1718
"name": "0.5.0 · 26.06",

src/megatron/bridge/data/datasets/gpt_sft.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -408,7 +408,7 @@ def _multiple_truncation(self, template_ids: list[list[int]], template_ids_keys:
408408

409409
if total_ids > self.max_seq_length:
410410
truncation_length_total = total_ids - self.max_seq_length
411-
num_fields = len(self.truncation_fields)
411+
num_fields = sum(key in self.truncation_fields for key in template_ids_keys)
412412
if num_fields > 0:
413413
# sorted equal divide length to each field
414414
# examples:

src/megatron/bridge/training/setup.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -596,6 +596,7 @@ def _build_distributed_model(cfg: ConfigContainer, pg_collection: ProcessGroupCo
596596
data_parallel_random_init=cfg.rng.data_parallel_random_init,
597597
)
598598
else:
599+
model_config.finalize()
599600
return model_config.provide_distributed_model(
600601
ddp_config=cfg.ddp,
601602
use_megatron_fsdp=cfg.dist.use_megatron_fsdp,

tests/unit_tests/data/datasets/test_gpt_sft.py

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -219,6 +219,25 @@ def test_multiple_truncation(self, tmp_path):
219219
assert context_ids == [101, 102, 103, 104, 201, 202, 203]
220220
assert label_ids == [301, 302]
221221

222+
def test_repeated_truncation_field_placeholder_handles_overflow(self, tmp_path):
223+
dataset_path = tmp_path / "repeated_prompt.jsonl"
224+
dataset_path.write_text(json.dumps({"input": "one two three four five", "output": "answer"}) + "\n")
225+
dataset = GPTSFTDataset(
226+
file_path=str(dataset_path),
227+
tokenizer=create_mock_tokenizer(),
228+
max_seq_length=8,
229+
max_num_samples=None,
230+
label_key="output",
231+
prompt_template="{input} Again: {input} {output}",
232+
truncation_field="input",
233+
memmap_workers=1,
234+
)
235+
236+
processed = dataset[0]
237+
238+
assert len(processed["input_ids"]) <= dataset.max_seq_length
239+
assert processed["answer_ids"]
240+
222241
def test_utils_func(self, tmp_path):
223242
dataset, _ = get_gpt_sft(tmp_path)
224243

tests/unit_tests/training/test_setup.py

Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,7 @@
1313
# limitations under the License.
1414

1515

16+
import json
1617
from types import SimpleNamespace
1718
from unittest.mock import MagicMock, Mock, patch
1819

@@ -25,6 +26,7 @@
2526
from megatron.bridge.models.gpt_provider import GPTModelProvider
2627
from megatron.bridge.models.hybrid.hybrid_builder import HybridModelConfig
2728
from megatron.bridge.models.hybrid.hybrid_provider import HybridModelProvider
29+
from megatron.bridge.models.llama_nemotron.llama_nemotron_provider import LlamaNemotronHeterogeneousProvider
2830
from megatron.bridge.models.transformer_config import TransformerConfig
2931
from megatron.bridge.training.callbacks import CallbackManager
3032
from megatron.bridge.training.checkpointing import load_checkpoint
@@ -556,9 +558,42 @@ def test_build_with_provider(self):
556558

557559
result = _build_distributed_model(cfg, pg_collection=MagicMock())
558560

561+
mock_provider.finalize.assert_called_once_with()
559562
mock_provider.provide_distributed_model.assert_called_once()
560563
assert result == mock_dist_model
561564

565+
def test_finalizes_heterogeneous_provider_before_entering_distributed_model_build(self):
566+
"""Finalize deferred heterogeneous fields before entering provider model construction."""
567+
block = {
568+
"attention": {"no_op": False, "replace_with_linear": False, "num_query_groups": 4},
569+
"mlp": {"no_op": False, "replace_with_linear": False, "ffn_hidden_size": 256},
570+
}
571+
provider = LlamaNemotronHeterogeneousProvider(
572+
num_layers=1,
573+
hidden_size=64,
574+
num_attention_heads=4,
575+
heterogeneous_layers_config_encoded_json=json.dumps({"block_configs": [block]}),
576+
)
577+
expected_model = [MagicMock()]
578+
579+
def provide_distributed_model(**_kwargs):
580+
assert len(provider.per_block_parameters) == 1
581+
return expected_model
582+
583+
provider.provide_distributed_model = Mock(side_effect=provide_distributed_model)
584+
cfg = SimpleNamespace(
585+
model=provider,
586+
ddp=MagicMock(),
587+
optimizer=SimpleNamespace(overlap_param_gather_with_optimizer_step=False),
588+
dist=SimpleNamespace(use_megatron_fsdp=False, use_torch_fsdp2=False),
589+
rng=SimpleNamespace(data_parallel_random_init=False),
590+
)
591+
592+
result = _build_distributed_model(cfg, pg_collection=MagicMock())
593+
594+
provider.provide_distributed_model.assert_called_once()
595+
assert result == expected_model
596+
562597

563598
def test_restart_rebinds_overlap_callbacks_to_rebuilt_model():
564599
"""A restart must replace callbacks bound to the discarded model."""

0 commit comments

Comments
 (0)