Skip to content

Commit 2366df8

Browse files
author
intern_nem_dev_2
committed
Propagate DataBlend revisions through pretrain SFT planning
1 parent 83119f9 commit 2366df8

10 files changed

Lines changed: 320 additions & 27 deletions

File tree

src/nemotron/data_prep/core/work_items.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -50,6 +50,9 @@ class DatasetWorkItem:
5050
sample: str | int | None
5151
sample_seed: int
5252

53+
# Source revision propagated from DataBlend.Dataset for deterministic HF discovery.
54+
revision: str | None = None
55+
5356
# Resolved tokenizer config (for plan creation)
5457
tokenizer_config: dict = field(default_factory=dict)
5558

@@ -95,6 +98,9 @@ class SftDatasetWorkItem:
9598
sample: str | int | None
9699
sample_seed: int
97100

101+
# Source revision propagated from DataBlend.Dataset for deterministic HF discovery.
102+
revision: str | None = None
103+
98104
# Resolved tokenizer config (for plan creation)
99105
tokenizer_config: dict = field(default_factory=dict)
100106

src/nemotron/data_prep/recipes/pretrain.py

Lines changed: 12 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -47,13 +47,12 @@
4747
import hashlib
4848
import json
4949
import time
50+
from collections.abc import Callable, Mapping
5051
from dataclasses import asdict, dataclass
5152
from pathlib import Path
52-
from typing import TYPE_CHECKING, Any, Mapping
53+
from typing import TYPE_CHECKING, Any
5354

5455
import cosmos_xenna.pipelines.v1 as pipelines_v1
55-
56-
from collections.abc import Callable
5756
from fsspec import AbstractFileSystem
5857

5958
from nemotron.data_prep.config import (
@@ -64,10 +63,11 @@
6463
ObservabilityConfig,
6564
TokenizerConfig,
6665
)
67-
from nemotron.data_prep.observability import pipeline_wandb_hook
68-
from nemotron.data_prep.utils.filesystem import ensure_dir, get_filesystem, write_json
6966
from nemotron.data_prep.core.finalize import scan_dataset_receipts
7067
from nemotron.data_prep.core.planning import PlanRequest, resolve_tokenizer, verify_binidx_output
68+
from nemotron.data_prep.core.work_items import DatasetWorkItem, ShardWorkItem
69+
from nemotron.data_prep.observability import pipeline_wandb_hook
70+
from nemotron.data_prep.recipes.execution_mode import ExecutionModeRequest, resolve_execution_mode
7171
from nemotron.data_prep.stages import (
7272
BinIdxTokenizationStage,
7373
BinIdxTokenizationStageConfig,
@@ -77,9 +77,8 @@
7777
PlanStage,
7878
PlanStageConfig,
7979
)
80+
from nemotron.data_prep.utils.filesystem import ensure_dir, get_filesystem, write_json
8081
from nemotron.data_prep.utils.hf_env import detect_hf_env_vars
81-
from nemotron.data_prep.core.work_items import DatasetWorkItem, ShardWorkItem
82-
from nemotron.data_prep.recipes.execution_mode import ExecutionModeRequest, resolve_execution_mode
8382

8483
if TYPE_CHECKING:
8584
from nemotron.data_prep.blend import DataBlend
@@ -122,6 +121,7 @@ def to_plan_request(self, item: DatasetWorkItem) -> PlanRequest:
122121
weight=item.weight,
123122
split=item.split,
124123
subset=item.subset,
124+
revision=item.revision,
125125
text_field=item.text_field,
126126
),
127127
num_shards=item.num_shards,
@@ -183,7 +183,7 @@ def _normalize_tokenizer(tokenizer: TokenizerConfig | Mapping[str, Any] | str) -
183183

184184

185185
def setup_pretrain_run(
186-
blend: "DataBlend",
186+
blend: DataBlend,
187187
output_dir: str | Path,
188188
tokenizer: TokenizerConfig | Mapping[str, Any] | str,
189189
*,
@@ -234,6 +234,7 @@ def setup_pretrain_run(
234234
"weight": d.weight,
235235
"split": d.split,
236236
"subset": d.subset,
237+
"revision": d.revision,
237238
"text_field": getattr(d, "text_field", None) or text_field_default,
238239
}
239240
for d in blend.datasets
@@ -270,6 +271,7 @@ def setup_pretrain_run(
270271
weight=d.weight,
271272
split=d.split,
272273
subset=d.subset,
274+
revision=d.revision,
273275
text_field=getattr(d, "text_field", None) or text_field_default,
274276
run_hash=run_hash,
275277
run_dir=run_dir,
@@ -297,7 +299,7 @@ def setup_pretrain_run(
297299

298300
def finalize_pretrain_run(
299301
context: PretrainRunContext,
300-
blend: "DataBlend",
302+
blend: DataBlend,
301303
output_dir: str | Path,
302304
) -> FormatResult:
303305
"""
@@ -360,7 +362,7 @@ def finalize_pretrain_run(
360362

361363

362364
def run_pretrain_pipeline(
363-
blend: "DataBlend",
365+
blend: DataBlend,
364366
output_dir: str | Path,
365367
tokenizer: TokenizerConfig | Mapping[str, Any] | str,
366368
*,

src/nemotron/data_prep/recipes/sft.py

Lines changed: 12 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -46,13 +46,12 @@
4646
import hashlib
4747
import json
4848
import time
49+
from collections.abc import Callable, Mapping
4950
from dataclasses import asdict, dataclass
5051
from pathlib import Path
51-
from typing import TYPE_CHECKING, Any, Mapping
52+
from typing import TYPE_CHECKING, Any
5253

5354
import cosmos_xenna.pipelines.v1 as pipelines_v1
54-
55-
from collections.abc import Callable
5655
from fsspec import AbstractFileSystem
5756

5857
from nemotron.data_prep.config import (
@@ -63,10 +62,11 @@
6362
ObservabilityConfig,
6463
TokenizerConfig,
6564
)
66-
from nemotron.data_prep.observability import pipeline_wandb_hook
67-
from nemotron.data_prep.utils.filesystem import ensure_dir, get_filesystem, write_json
6865
from nemotron.data_prep.core.finalize import scan_dataset_receipts
6966
from nemotron.data_prep.core.planning import PlanRequest, resolve_tokenizer, verify_parquet_output
67+
from nemotron.data_prep.core.work_items import SftDatasetWorkItem, SftShardWorkItem
68+
from nemotron.data_prep.observability import pipeline_wandb_hook
69+
from nemotron.data_prep.recipes.execution_mode import ExecutionModeRequest, resolve_execution_mode
7070
from nemotron.data_prep.stages import (
7171
DownloadStage,
7272
DownloadStageConfig,
@@ -76,9 +76,8 @@
7676
PlanStage,
7777
SftPlanStageConfig,
7878
)
79+
from nemotron.data_prep.utils.filesystem import ensure_dir, get_filesystem, write_json
7980
from nemotron.data_prep.utils.hf_env import detect_hf_env_vars
80-
from nemotron.data_prep.core.work_items import SftDatasetWorkItem, SftShardWorkItem
81-
from nemotron.data_prep.recipes.execution_mode import ExecutionModeRequest, resolve_execution_mode
8281

8382
if TYPE_CHECKING:
8483
from nemotron.data_prep.blend import DataBlend
@@ -121,6 +120,7 @@ def to_plan_request(self, item: SftDatasetWorkItem) -> PlanRequest:
121120
weight=item.weight,
122121
split=item.split,
123122
subset=item.subset,
123+
revision=item.revision,
124124
text_field=item.messages_field,
125125
),
126126
num_shards=item.num_shards,
@@ -190,7 +190,7 @@ def _normalize_tokenizer(tokenizer: TokenizerConfig | Mapping[str, Any] | str) -
190190

191191

192192
def setup_sft_run(
193-
blend: "DataBlend",
193+
blend: DataBlend,
194194
output_dir: str | Path,
195195
tokenizer: TokenizerConfig | Mapping[str, Any] | str,
196196
*,
@@ -257,6 +257,7 @@ def setup_sft_run(
257257
"weight": d.weight,
258258
"split": d.split,
259259
"subset": d.subset,
260+
"revision": d.revision,
260261
"messages_field": getattr(d, "messages_field", None) or messages_field_default,
261262
"tools_field": getattr(d, "tools_field", None) or tools_field_default,
262263
}
@@ -304,6 +305,7 @@ def setup_sft_run(
304305
weight=d.weight,
305306
split=d.split,
306307
subset=d.subset,
308+
revision=d.revision,
307309
run_hash=run_hash,
308310
run_dir=run_dir,
309311
config_hash=config_hash,
@@ -340,7 +342,7 @@ def setup_sft_run(
340342

341343
def finalize_sft_run(
342344
context: SftRunContext,
343-
blend: "DataBlend",
345+
blend: DataBlend,
344346
output_dir: str | Path,
345347
) -> FormatResult:
346348
"""
@@ -403,7 +405,7 @@ def finalize_sft_run(
403405

404406

405407
def run_sft_pipeline(
406-
blend: "DataBlend",
408+
blend: DataBlend,
407409
output_dir: str | Path,
408410
tokenizer: TokenizerConfig | Mapping[str, Any] | str,
409411
*,

src/nemotron/kit/artifacts/pretrain_blends.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -104,15 +104,15 @@ def get_input_uris(self) -> list[str]:
104104
@classmethod
105105
def from_result(
106106
cls,
107-
format_result: "FormatResult",
108-
blend: "DataBlend",
107+
format_result: FormatResult,
108+
blend: DataBlend,
109109
tokenizer_model: str,
110110
blend_json_path: str | Path,
111111
*,
112112
text_field_default: str = "text",
113113
elapsed_sec: float = 0.0,
114114
name: str | None = None,
115-
) -> "PretrainBlendsArtifact":
115+
) -> PretrainBlendsArtifact:
116116
"""Create artifact from pipeline format result.
117117
118118
This is a convenience constructor that builds the source_datasets
@@ -137,6 +137,7 @@ def from_result(
137137
weight=d.weight,
138138
split=d.split,
139139
subset=d.subset,
140+
revision=d.revision,
140141
text_field=d.text_field or text_field_default,
141142
)
142143
for d in blend.datasets

src/nemotron/kit/artifacts/sft_data.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -174,6 +174,7 @@ def from_result(
174174
weight=d.weight,
175175
split=d.split,
176176
subset=d.subset,
177+
revision=d.revision,
177178
text_field=d.text_field or messages_field_default,
178179
)
179180
for d in blend.datasets

0 commit comments

Comments
 (0)