ftrain 9.0.1__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- ftrain/__init__.py +25 -0
- ftrain/api.py +605 -0
- ftrain/callbacks.py +954 -0
- ftrain/captain.py +2207 -0
- ftrain/config.py +1659 -0
- ftrain/core.py +3885 -0
- ftrain/cpp_merge.py +995 -0
- ftrain/dashboard.py +1300 -0
- ftrain/data_quality.py +1156 -0
- ftrain/data_utils.py +646 -0
- ftrain/dataset.py +1034 -0
- ftrain/families.py +530 -0
- ftrain/kernels_dora.py +767 -0
- ftrain/lora.py +877 -0
- ftrain/lora_dora.py +1160 -0
- ftrain/merge_advanced.py +1465 -0
- ftrain/merge_intel.py +1766 -0
- ftrain/merger.py +2079 -0
- ftrain/model_utils.py +1126 -0
- ftrain/projection.py +735 -0
- ftrain/rewards.py +1875 -0
- ftrain/safety.py +574 -0
- ftrain/similarity.py +1183 -0
- ftrain/speed.py +26 -0
- ftrain/tensor_stats.py +57 -0
- ftrain/train_optim.py +66 -0
- ftrain/ui.py +1125 -0
- ftrain-9.0.1.dist-info/METADATA +22 -0
- ftrain-9.0.1.dist-info/RECORD +32 -0
- ftrain-9.0.1.dist-info/WHEEL +5 -0
- ftrain-9.0.1.dist-info/licenses/LICENSE +201 -0
- ftrain-9.0.1.dist-info/top_level.txt +1 -0
ftrain/__init__.py
ADDED
|
@@ -0,0 +1,25 @@
|
|
|
1
|
+
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
import logging
|
|
5
|
+
from typing import Final
|
|
6
|
+
|
|
7
|
+
from .api import merge, test, train
|
|
8
|
+
|
|
9
|
+
__version__: Final[str] = "1.1.0"
|
|
10
|
+
__author__: Final[str] = "FTRAIN Engine Team"
|
|
11
|
+
|
|
12
|
+
# Public package surface.
|
|
13
|
+
__all__: Final[tuple[str, ...]] = (
|
|
14
|
+
"train",
|
|
15
|
+
"merge",
|
|
16
|
+
"test",
|
|
17
|
+
"__version__",
|
|
18
|
+
"__author__",
|
|
19
|
+
)
|
|
20
|
+
|
|
21
|
+
_logger = logging.getLogger(__name__)
|
|
22
|
+
|
|
23
|
+
if not any(isinstance(handler, logging.NullHandler) for handler in _logger.handlers):
|
|
24
|
+
_logger.addHandler(logging.NullHandler())
|
|
25
|
+
_logger.propagate = True
|
ftrain/api.py
ADDED
|
@@ -0,0 +1,605 @@
|
|
|
1
|
+
"""
|
|
2
|
+
FTRAIN High-Level API
|
|
3
|
+
=====================
|
|
4
|
+
|
|
5
|
+
Stable public entry points for the FTRAIN training, merging, and diagnostic
|
|
6
|
+
pipelines.
|
|
7
|
+
|
|
8
|
+
Public API:
|
|
9
|
+
train.fire(...)
|
|
10
|
+
merge.fire(...)
|
|
11
|
+
test()
|
|
12
|
+
|
|
13
|
+
The implementation in this module intentionally acts as an orchestration
|
|
14
|
+
layer. Heavy model/training/merging logic belongs in the underlying modules
|
|
15
|
+
(`core`, `merger`, `data_utils`, etc.).
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
from __future__ import annotations
|
|
19
|
+
|
|
20
|
+
import logging
|
|
21
|
+
from typing import Any, Dict, Mapping, Optional, Tuple
|
|
22
|
+
|
|
23
|
+
from . import rewards
|
|
24
|
+
from .config import MergeConfig, TrainConfig
|
|
25
|
+
from .core import Ftrain
|
|
26
|
+
from .data_utils import load_data
|
|
27
|
+
from .merger import Merger
|
|
28
|
+
|
|
29
|
+
__all__ = [
|
|
30
|
+
"train",
|
|
31
|
+
"merge",
|
|
32
|
+
"test",
|
|
33
|
+
"xml_format_reward",
|
|
34
|
+
"math_exact_reward",
|
|
35
|
+
"python_exec_reward",
|
|
36
|
+
]
|
|
37
|
+
|
|
38
|
+
LOGGER = logging.getLogger(__name__)
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
# ---------------------------------------------------------------------------
|
|
42
|
+
# Internal helpers
|
|
43
|
+
# ---------------------------------------------------------------------------
|
|
44
|
+
|
|
45
|
+
_DEFAULT_TRAIN_OUTPUT_DIR = "./ftrain_output"
|
|
46
|
+
_DEFAULT_MERGE_OUTPUT_DIR = "./merged_model"
|
|
47
|
+
|
|
48
|
+
_MIN_VALIDATION_RATIO = 0.0
|
|
49
|
+
_MAX_VALIDATION_RATIO = 1.0
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def _validate_model_name(model: Optional[str], argument_name: str) -> str:
|
|
53
|
+
"""
|
|
54
|
+
Validate a model identifier/path before passing it deeper into FTRAIN.
|
|
55
|
+
|
|
56
|
+
Model names may be Hugging Face IDs, local paths, or other identifiers,
|
|
57
|
+
so this helper intentionally does not try to validate filesystem or
|
|
58
|
+
repository existence here.
|
|
59
|
+
"""
|
|
60
|
+
if model is None:
|
|
61
|
+
raise ValueError(f"'{argument_name}' must be provided.")
|
|
62
|
+
|
|
63
|
+
if not isinstance(model, str):
|
|
64
|
+
raise TypeError(
|
|
65
|
+
f"'{argument_name}' must be a string, got {type(model).__name__}."
|
|
66
|
+
)
|
|
67
|
+
|
|
68
|
+
model = model.strip()
|
|
69
|
+
|
|
70
|
+
if not model:
|
|
71
|
+
raise ValueError(f"'{argument_name}' cannot be empty.")
|
|
72
|
+
|
|
73
|
+
return model
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
def _validate_steps(steps: int) -> int:
|
|
77
|
+
"""Validate the requested number of training steps."""
|
|
78
|
+
if isinstance(steps, bool) or not isinstance(steps, int):
|
|
79
|
+
raise TypeError(
|
|
80
|
+
f"'Steps' must be an integer, got {type(steps).__name__}."
|
|
81
|
+
)
|
|
82
|
+
|
|
83
|
+
if steps <= 0:
|
|
84
|
+
raise ValueError(f"'Steps' must be greater than zero, got {steps}.")
|
|
85
|
+
|
|
86
|
+
return steps
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
def _validate_validation_ratio(ratio: float) -> float:
|
|
90
|
+
"""Validate and normalize a dataset validation split ratio."""
|
|
91
|
+
if isinstance(ratio, bool) or not isinstance(ratio, (int, float)):
|
|
92
|
+
raise TypeError(
|
|
93
|
+
"'validation_ratio' must be a number between 0 and 1, "
|
|
94
|
+
f"got {type(ratio).__name__}."
|
|
95
|
+
)
|
|
96
|
+
|
|
97
|
+
ratio = float(ratio)
|
|
98
|
+
|
|
99
|
+
if not _MIN_VALIDATION_RATIO <= ratio <= _MAX_VALIDATION_RATIO:
|
|
100
|
+
raise ValueError(
|
|
101
|
+
"'validation_ratio' must be between 0.0 and 1.0, "
|
|
102
|
+
f"got {ratio}."
|
|
103
|
+
)
|
|
104
|
+
|
|
105
|
+
return ratio
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def _split_data(
|
|
109
|
+
data: Any,
|
|
110
|
+
validation_ratio: float = 0.10,
|
|
111
|
+
) -> Tuple[Any, Any]:
|
|
112
|
+
"""
|
|
113
|
+
Split loaded data into training and validation portions.
|
|
114
|
+
|
|
115
|
+
The function intentionally operates on the object returned by
|
|
116
|
+
``load_data`` rather than assuming it is a Python list. This keeps it
|
|
117
|
+
compatible with sequence-like datasets and common dataset abstractions.
|
|
118
|
+
|
|
119
|
+
For a non-empty dataset:
|
|
120
|
+
- training data always receives at least one sample
|
|
121
|
+
- validation data is ``None`` when a validation split is impossible
|
|
122
|
+
- validation ratio ``0`` disables validation
|
|
123
|
+
|
|
124
|
+
Returns:
|
|
125
|
+
(train_data, validation_data)
|
|
126
|
+
"""
|
|
127
|
+
validation_ratio = _validate_validation_ratio(validation_ratio)
|
|
128
|
+
|
|
129
|
+
try:
|
|
130
|
+
dataset_size = len(data)
|
|
131
|
+
except TypeError as exc:
|
|
132
|
+
raise TypeError(
|
|
133
|
+
"The object returned by 'load_data()' must be sized "
|
|
134
|
+
"(it must implement __len__)."
|
|
135
|
+
) from exc
|
|
136
|
+
|
|
137
|
+
if dataset_size < 0:
|
|
138
|
+
raise ValueError(
|
|
139
|
+
f"Loaded dataset reported an invalid length: {dataset_size}."
|
|
140
|
+
)
|
|
141
|
+
|
|
142
|
+
if dataset_size == 0:
|
|
143
|
+
raise ValueError(
|
|
144
|
+
"The loaded dataset is empty. FTRAIN cannot start training "
|
|
145
|
+
"without at least one training sample."
|
|
146
|
+
)
|
|
147
|
+
|
|
148
|
+
if dataset_size == 1 or validation_ratio <= 0.0:
|
|
149
|
+
return data, None
|
|
150
|
+
|
|
151
|
+
# Calculate validation count while guaranteeing at least one training
|
|
152
|
+
# sample. This fixes the original behavior where small datasets could
|
|
153
|
+
# accidentally produce an empty training split.
|
|
154
|
+
validation_size = int(round(dataset_size * validation_ratio))
|
|
155
|
+
validation_size = max(1, validation_size)
|
|
156
|
+
validation_size = min(validation_size, dataset_size - 1)
|
|
157
|
+
|
|
158
|
+
split_index = dataset_size - validation_size
|
|
159
|
+
|
|
160
|
+
try:
|
|
161
|
+
train_data = data[:split_index]
|
|
162
|
+
validation_data = data[split_index:]
|
|
163
|
+
except (TypeError, IndexError) as exc:
|
|
164
|
+
raise TypeError(
|
|
165
|
+
"The loaded dataset does not support slicing. "
|
|
166
|
+
"FTRAIN currently requires a sliceable dataset for its "
|
|
167
|
+
"automatic train/validation split."
|
|
168
|
+
) from exc
|
|
169
|
+
|
|
170
|
+
if len(train_data) == 0:
|
|
171
|
+
raise RuntimeError(
|
|
172
|
+
"Internal dataset splitting error: training split is empty."
|
|
173
|
+
)
|
|
174
|
+
|
|
175
|
+
if len(validation_data) == 0:
|
|
176
|
+
LOGGER.warning(
|
|
177
|
+
"Validation splitting produced an empty validation set; "
|
|
178
|
+
"continuing without validation."
|
|
179
|
+
)
|
|
180
|
+
validation_data = None
|
|
181
|
+
|
|
182
|
+
return train_data, validation_data
|
|
183
|
+
|
|
184
|
+
|
|
185
|
+
def _build_train_config(
|
|
186
|
+
*,
|
|
187
|
+
model: str,
|
|
188
|
+
captain: Optional[str],
|
|
189
|
+
steps: int,
|
|
190
|
+
answer: str,
|
|
191
|
+
output_dir: str,
|
|
192
|
+
extra_kwargs: Mapping[str, Any],
|
|
193
|
+
) -> TrainConfig:
|
|
194
|
+
"""
|
|
195
|
+
Construct TrainConfig while preventing duplicate keyword collisions.
|
|
196
|
+
|
|
197
|
+
Explicit public API arguments always take precedence over values supplied
|
|
198
|
+
through **kwargs.
|
|
199
|
+
"""
|
|
200
|
+
config_kwargs: Dict[str, Any] = dict(extra_kwargs)
|
|
201
|
+
|
|
202
|
+
# These are explicitly controlled by the public API.
|
|
203
|
+
reserved = {
|
|
204
|
+
"model_name",
|
|
205
|
+
"captain_model",
|
|
206
|
+
"max_steps",
|
|
207
|
+
"answer_mode",
|
|
208
|
+
"captain_mode",
|
|
209
|
+
"output_dir",
|
|
210
|
+
}
|
|
211
|
+
|
|
212
|
+
for key in reserved:
|
|
213
|
+
config_kwargs.pop(key, None)
|
|
214
|
+
|
|
215
|
+
config_kwargs.update(
|
|
216
|
+
{
|
|
217
|
+
"model_name": model,
|
|
218
|
+
"captain_model": captain,
|
|
219
|
+
"max_steps": steps,
|
|
220
|
+
"answer_mode": answer,
|
|
221
|
+
"captain_mode": "llm" if captain else "rule",
|
|
222
|
+
"output_dir": output_dir,
|
|
223
|
+
}
|
|
224
|
+
)
|
|
225
|
+
|
|
226
|
+
try:
|
|
227
|
+
return TrainConfig(**config_kwargs)
|
|
228
|
+
except TypeError as exc:
|
|
229
|
+
raise TypeError(
|
|
230
|
+
"Failed to construct TrainConfig. "
|
|
231
|
+
"Check the supplied training options and make sure every "
|
|
232
|
+
"keyword is supported by your installed TrainConfig."
|
|
233
|
+
) from exc
|
|
234
|
+
|
|
235
|
+
|
|
236
|
+
def _build_merge_config(
|
|
237
|
+
*,
|
|
238
|
+
model_a: str,
|
|
239
|
+
model_b: str,
|
|
240
|
+
captain: Optional[str],
|
|
241
|
+
output_dir: str,
|
|
242
|
+
extra_kwargs: Mapping[str, Any],
|
|
243
|
+
) -> MergeConfig:
|
|
244
|
+
"""
|
|
245
|
+
Construct MergeConfig while preventing duplicate keyword collisions.
|
|
246
|
+
"""
|
|
247
|
+
config_kwargs: Dict[str, Any] = dict(extra_kwargs)
|
|
248
|
+
|
|
249
|
+
reserved = {
|
|
250
|
+
"model_a",
|
|
251
|
+
"model_b",
|
|
252
|
+
"captain_model",
|
|
253
|
+
"output_dir",
|
|
254
|
+
}
|
|
255
|
+
|
|
256
|
+
for key in reserved:
|
|
257
|
+
config_kwargs.pop(key, None)
|
|
258
|
+
|
|
259
|
+
config_kwargs.update(
|
|
260
|
+
{
|
|
261
|
+
"model_a": model_a,
|
|
262
|
+
"model_b": model_b,
|
|
263
|
+
"captain_model": captain,
|
|
264
|
+
"output_dir": output_dir,
|
|
265
|
+
}
|
|
266
|
+
)
|
|
267
|
+
|
|
268
|
+
try:
|
|
269
|
+
return MergeConfig(**config_kwargs)
|
|
270
|
+
except TypeError as exc:
|
|
271
|
+
raise TypeError(
|
|
272
|
+
"Failed to construct MergeConfig. "
|
|
273
|
+
"Check the supplied merge options and make sure every keyword "
|
|
274
|
+
"is supported by your installed MergeConfig."
|
|
275
|
+
) from exc
|
|
276
|
+
|
|
277
|
+
|
|
278
|
+
# ---------------------------------------------------------------------------
|
|
279
|
+
# Training API
|
|
280
|
+
# ---------------------------------------------------------------------------
|
|
281
|
+
|
|
282
|
+
|
|
283
|
+
class train:
|
|
284
|
+
"""
|
|
285
|
+
High-level FTRAIN training interface.
|
|
286
|
+
|
|
287
|
+
Example:
|
|
288
|
+
train.fire(
|
|
289
|
+
Model="Qwen/Qwen2.5-0.5B-Instruct",
|
|
290
|
+
Data="dataset.json",
|
|
291
|
+
Steps=500,
|
|
292
|
+
)
|
|
293
|
+
"""
|
|
294
|
+
|
|
295
|
+
@staticmethod
|
|
296
|
+
def fire(
|
|
297
|
+
Model: str,
|
|
298
|
+
Data: Any,
|
|
299
|
+
Steps: int = 100,
|
|
300
|
+
Captain: Optional[str] = None,
|
|
301
|
+
Answer: str = "auto_yes",
|
|
302
|
+
**kwargs: Any,
|
|
303
|
+
) -> Any:
|
|
304
|
+
"""
|
|
305
|
+
Run the complete FTRAIN training pipeline.
|
|
306
|
+
|
|
307
|
+
Parameters:
|
|
308
|
+
Model:
|
|
309
|
+
Base model identifier or local model path.
|
|
310
|
+
|
|
311
|
+
Data:
|
|
312
|
+
Dataset accepted by ``load_data``.
|
|
313
|
+
|
|
314
|
+
Steps:
|
|
315
|
+
Maximum number of training steps.
|
|
316
|
+
|
|
317
|
+
Captain:
|
|
318
|
+
Optional captain/reviewer LLM model.
|
|
319
|
+
|
|
320
|
+
Answer:
|
|
321
|
+
Answer/reward mode.
|
|
322
|
+
|
|
323
|
+
output_dir:
|
|
324
|
+
Output directory. Defaults to ``./ftrain_output``.
|
|
325
|
+
|
|
326
|
+
validation_ratio:
|
|
327
|
+
Fraction of the dataset reserved for validation.
|
|
328
|
+
Defaults to ``0.10``.
|
|
329
|
+
|
|
330
|
+
Returns:
|
|
331
|
+
Whatever ``Ftrain(...).train()`` returns.
|
|
332
|
+
|
|
333
|
+
Raises:
|
|
334
|
+
ValueError:
|
|
335
|
+
Invalid required arguments or empty dataset.
|
|
336
|
+
|
|
337
|
+
TypeError:
|
|
338
|
+
Invalid argument types or unsupported configuration options.
|
|
339
|
+
"""
|
|
340
|
+
model = _validate_model_name(Model, "Model")
|
|
341
|
+
steps = _validate_steps(Steps)
|
|
342
|
+
|
|
343
|
+
if not isinstance(Answer, str):
|
|
344
|
+
raise TypeError(
|
|
345
|
+
f"'Answer' must be a string, got {type(Answer).__name__}."
|
|
346
|
+
)
|
|
347
|
+
|
|
348
|
+
answer = Answer.strip()
|
|
349
|
+
|
|
350
|
+
if not answer:
|
|
351
|
+
raise ValueError("'Answer' cannot be empty.")
|
|
352
|
+
|
|
353
|
+
if Captain is not None:
|
|
354
|
+
captain = _validate_model_name(Captain, "Captain")
|
|
355
|
+
else:
|
|
356
|
+
captain = None
|
|
357
|
+
|
|
358
|
+
runtime_kwargs: Dict[str, Any] = dict(kwargs)
|
|
359
|
+
|
|
360
|
+
output_dir = runtime_kwargs.pop(
|
|
361
|
+
"output_dir",
|
|
362
|
+
_DEFAULT_TRAIN_OUTPUT_DIR,
|
|
363
|
+
)
|
|
364
|
+
|
|
365
|
+
if output_dir is None:
|
|
366
|
+
output_dir = _DEFAULT_TRAIN_OUTPUT_DIR
|
|
367
|
+
|
|
368
|
+
if not isinstance(output_dir, str):
|
|
369
|
+
raise TypeError(
|
|
370
|
+
"'output_dir' must be a string, "
|
|
371
|
+
f"got {type(output_dir).__name__}."
|
|
372
|
+
)
|
|
373
|
+
|
|
374
|
+
output_dir = output_dir.strip()
|
|
375
|
+
|
|
376
|
+
if not output_dir:
|
|
377
|
+
raise ValueError("'output_dir' cannot be empty.")
|
|
378
|
+
|
|
379
|
+
# Support an explicit validation ratio without leaking this
|
|
380
|
+
# orchestration-level setting into TrainConfig unless that config
|
|
381
|
+
# explicitly expects it.
|
|
382
|
+
validation_ratio = runtime_kwargs.pop(
|
|
383
|
+
"validation_ratio",
|
|
384
|
+
runtime_kwargs.pop("val_split", 0.10),
|
|
385
|
+
)
|
|
386
|
+
|
|
387
|
+
LOGGER.info(
|
|
388
|
+
"Starting FTRAIN training: model=%s, steps=%d, captain=%s, "
|
|
389
|
+
"output_dir=%s",
|
|
390
|
+
model,
|
|
391
|
+
steps,
|
|
392
|
+
captain or "disabled",
|
|
393
|
+
output_dir,
|
|
394
|
+
)
|
|
395
|
+
|
|
396
|
+
# Load data once. Any expensive parsing/tokenization handled by
|
|
397
|
+
# load_data therefore remains centralized.
|
|
398
|
+
data = load_data(Data)
|
|
399
|
+
|
|
400
|
+
train_data, val_data = _split_data(
|
|
401
|
+
data,
|
|
402
|
+
validation_ratio=validation_ratio,
|
|
403
|
+
)
|
|
404
|
+
|
|
405
|
+
LOGGER.info(
|
|
406
|
+
"Dataset prepared: total=%d, train=%d, validation=%s",
|
|
407
|
+
len(data),
|
|
408
|
+
len(train_data),
|
|
409
|
+
len(val_data) if val_data is not None else "disabled",
|
|
410
|
+
)
|
|
411
|
+
|
|
412
|
+
config = _build_train_config(
|
|
413
|
+
model=model,
|
|
414
|
+
captain=captain,
|
|
415
|
+
steps=steps,
|
|
416
|
+
answer=answer,
|
|
417
|
+
output_dir=output_dir,
|
|
418
|
+
extra_kwargs=runtime_kwargs,
|
|
419
|
+
)
|
|
420
|
+
|
|
421
|
+
engine = Ftrain(
|
|
422
|
+
config,
|
|
423
|
+
train_data,
|
|
424
|
+
val_data,
|
|
425
|
+
)
|
|
426
|
+
|
|
427
|
+
result = engine.train()
|
|
428
|
+
|
|
429
|
+
LOGGER.info(
|
|
430
|
+
"FTRAIN training pipeline completed successfully for model=%s",
|
|
431
|
+
model,
|
|
432
|
+
)
|
|
433
|
+
|
|
434
|
+
return result
|
|
435
|
+
|
|
436
|
+
|
|
437
|
+
# ---------------------------------------------------------------------------
|
|
438
|
+
# Model merge API
|
|
439
|
+
# ---------------------------------------------------------------------------
|
|
440
|
+
|
|
441
|
+
|
|
442
|
+
class merge:
|
|
443
|
+
"""
|
|
444
|
+
High-level FTRAIN model merging interface.
|
|
445
|
+
|
|
446
|
+
Example:
|
|
447
|
+
merge.fire(
|
|
448
|
+
First="model_a",
|
|
449
|
+
Second="model_b",
|
|
450
|
+
output_dir="./merged",
|
|
451
|
+
)
|
|
452
|
+
"""
|
|
453
|
+
|
|
454
|
+
@staticmethod
|
|
455
|
+
def fire(
|
|
456
|
+
First: Optional[str] = None,
|
|
457
|
+
Second: Optional[str] = None,
|
|
458
|
+
Captain: Optional[str] = None,
|
|
459
|
+
**kwargs: Any,
|
|
460
|
+
) -> Any:
|
|
461
|
+
"""
|
|
462
|
+
Run the FTRAIN intelligent model-merging pipeline.
|
|
463
|
+
|
|
464
|
+
Both ``First`` and ``Second`` are required.
|
|
465
|
+
|
|
466
|
+
Backward-compatible aliases are also accepted:
|
|
467
|
+
Model_a
|
|
468
|
+
Model_b
|
|
469
|
+
|
|
470
|
+
Explicit ``First``/``Second`` values take precedence over aliases.
|
|
471
|
+
"""
|
|
472
|
+
runtime_kwargs: Dict[str, Any] = dict(kwargs)
|
|
473
|
+
|
|
474
|
+
# Preserve compatibility with the previous API.
|
|
475
|
+
model_a = First
|
|
476
|
+
model_b = Second
|
|
477
|
+
|
|
478
|
+
if model_a is None:
|
|
479
|
+
model_a = runtime_kwargs.pop("Model_a", None)
|
|
480
|
+
else:
|
|
481
|
+
# Do not allow a stale alias to create confusing behavior.
|
|
482
|
+
runtime_kwargs.pop("Model_a", None)
|
|
483
|
+
|
|
484
|
+
if model_b is None:
|
|
485
|
+
model_b = runtime_kwargs.pop("Model_b", None)
|
|
486
|
+
else:
|
|
487
|
+
runtime_kwargs.pop("Model_b", None)
|
|
488
|
+
|
|
489
|
+
model_a = _validate_model_name(model_a, "First")
|
|
490
|
+
model_b = _validate_model_name(model_b, "Second")
|
|
491
|
+
|
|
492
|
+
if Captain is not None:
|
|
493
|
+
captain = _validate_model_name(Captain, "Captain")
|
|
494
|
+
else:
|
|
495
|
+
captain = None
|
|
496
|
+
|
|
497
|
+
# Avoid accidentally merging a model with itself unless the caller
|
|
498
|
+
# explicitly disables the check via allow_self_merge=True.
|
|
499
|
+
allow_self_merge = bool(
|
|
500
|
+
runtime_kwargs.pop("allow_self_merge", False)
|
|
501
|
+
)
|
|
502
|
+
|
|
503
|
+
if model_a == model_b and not allow_self_merge:
|
|
504
|
+
raise ValueError(
|
|
505
|
+
"The two merge inputs resolve to the same model. "
|
|
506
|
+
"Pass allow_self_merge=True only when an intentional "
|
|
507
|
+
"self-merge is required."
|
|
508
|
+
)
|
|
509
|
+
|
|
510
|
+
output_dir = runtime_kwargs.pop(
|
|
511
|
+
"output_dir",
|
|
512
|
+
_DEFAULT_MERGE_OUTPUT_DIR,
|
|
513
|
+
)
|
|
514
|
+
|
|
515
|
+
if output_dir is None:
|
|
516
|
+
output_dir = _DEFAULT_MERGE_OUTPUT_DIR
|
|
517
|
+
|
|
518
|
+
if not isinstance(output_dir, str):
|
|
519
|
+
raise TypeError(
|
|
520
|
+
"'output_dir' must be a string, "
|
|
521
|
+
f"got {type(output_dir).__name__}."
|
|
522
|
+
)
|
|
523
|
+
|
|
524
|
+
output_dir = output_dir.strip()
|
|
525
|
+
|
|
526
|
+
if not output_dir:
|
|
527
|
+
raise ValueError("'output_dir' cannot be empty.")
|
|
528
|
+
|
|
529
|
+
LOGGER.info(
|
|
530
|
+
"Starting FTRAIN merge: model_a=%s, model_b=%s, captain=%s, "
|
|
531
|
+
"output_dir=%s",
|
|
532
|
+
model_a,
|
|
533
|
+
model_b,
|
|
534
|
+
captain or "disabled",
|
|
535
|
+
output_dir,
|
|
536
|
+
)
|
|
537
|
+
|
|
538
|
+
config = _build_merge_config(
|
|
539
|
+
model_a=model_a,
|
|
540
|
+
model_b=model_b,
|
|
541
|
+
captain=captain,
|
|
542
|
+
output_dir=output_dir,
|
|
543
|
+
extra_kwargs=runtime_kwargs,
|
|
544
|
+
)
|
|
545
|
+
|
|
546
|
+
merger = Merger(config)
|
|
547
|
+
result = merger.merge()
|
|
548
|
+
|
|
549
|
+
LOGGER.info(
|
|
550
|
+
"FTRAIN model merge completed successfully: output_dir=%s",
|
|
551
|
+
output_dir,
|
|
552
|
+
)
|
|
553
|
+
|
|
554
|
+
return result
|
|
555
|
+
|
|
556
|
+
|
|
557
|
+
# ---------------------------------------------------------------------------
|
|
558
|
+
# Diagnostics / smoke test
|
|
559
|
+
# ---------------------------------------------------------------------------
|
|
560
|
+
|
|
561
|
+
|
|
562
|
+
def test() -> bool:
|
|
563
|
+
"""
|
|
564
|
+
Run a lightweight package-level health check.
|
|
565
|
+
|
|
566
|
+
This intentionally does not load a model or allocate GPU memory.
|
|
567
|
+
It verifies that the public API and reward exports are available.
|
|
568
|
+
|
|
569
|
+
Returns:
|
|
570
|
+
``True`` when the package API is available.
|
|
571
|
+
|
|
572
|
+
Raises:
|
|
573
|
+
RuntimeError:
|
|
574
|
+
If a required public component is unexpectedly unavailable.
|
|
575
|
+
"""
|
|
576
|
+
required_objects = {
|
|
577
|
+
"train.fire": getattr(train, "fire", None),
|
|
578
|
+
"merge.fire": getattr(merge, "fire", None),
|
|
579
|
+
"xml_format_reward": xml_format_reward,
|
|
580
|
+
"math_exact_reward": math_exact_reward,
|
|
581
|
+
"python_exec_reward": python_exec_reward,
|
|
582
|
+
}
|
|
583
|
+
|
|
584
|
+
missing = [
|
|
585
|
+
name
|
|
586
|
+
for name, obj in required_objects.items()
|
|
587
|
+
if obj is None or not callable(obj)
|
|
588
|
+
]
|
|
589
|
+
|
|
590
|
+
if missing:
|
|
591
|
+
raise RuntimeError(
|
|
592
|
+
"FTRAIN API health check failed. Missing or invalid exports: "
|
|
593
|
+
+ ", ".join(missing)
|
|
594
|
+
)
|
|
595
|
+
|
|
596
|
+
print("✅ FTRAIN API health check passed")
|
|
597
|
+
print("✅ train.fire available")
|
|
598
|
+
print("✅ merge.fire available")
|
|
599
|
+
print("✅ reward functions available")
|
|
600
|
+
|
|
601
|
+
return True
|
|
602
|
+
|
|
603
|
+
xml_format_reward = rewards.xml_format_reward
|
|
604
|
+
math_exact_reward = rewards.math_exact_reward
|
|
605
|
+
python_exec_reward = rewards.python_exec_reward
|