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 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