-
Notifications
You must be signed in to change notification settings - Fork 7
Expand file tree
/
Copy pathtrain.py
More file actions
133 lines (110 loc) · 3.91 KB
/
Copy pathtrain.py
File metadata and controls
133 lines (110 loc) · 3.91 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
"""``studentsim-train``: launch Stage 1 or Stage 2 SFT for one domain.
Usage
-----
Stage 1 (pooled training)::
studentsim-train --config configs/training/stage1_chess.yaml
Stage 2 over a roster of students::
studentsim-train --config configs/training/stage2_chess.yaml \\
--roster data/chess/rosters/stage2.json
Stage 2 for one specific student::
studentsim-train --config configs/training/stage2_chess.yaml --student-id alice
The ``data_path`` / ``output_dir`` fields in Stage-2 YAMLs may contain a
``{student_id}`` placeholder which is substituted per-student.
"""
from __future__ import annotations
import argparse
import dataclasses
import json
import sys
from collections.abc import Sequence
from pathlib import Path
from studentsim.training import (
PerStudentDriver,
Stage1Trainer,
Stage2Trainer,
TrainingConfig,
)
def _substitute_student(text: str, student_id: str) -> str:
return text.replace("{student_id}", student_id)
def _per_student_config(base: TrainingConfig, student_id: str) -> TrainingConfig:
"""Substitute {student_id} placeholders in path-like fields."""
return dataclasses.replace(
base,
data_path=_substitute_student(base.data_path, student_id),
output_dir=_substitute_student(base.output_dir, student_id),
)
def main(argv: Sequence[str] | None = None) -> int:
parser = argparse.ArgumentParser(
prog="studentsim-train",
description="Stage 1 or Stage 2 SFT.",
)
parser.add_argument("--config", type=Path, required=True, help="Path to a training YAML.")
group = parser.add_mutually_exclusive_group()
group.add_argument(
"--student-id",
default=None,
help="Stage 2 single-student mode.",
)
group.add_argument(
"--roster",
type=Path,
default=None,
help="JSON list of student ids; runs Stage 2 for each in order.",
)
parser.add_argument(
"--fail-fast",
action="store_true",
help="With --roster, abort on the first per-student failure.",
)
parser.add_argument(
"--trainer-seed",
type=int,
default=None,
help=(
"Override TrainingConfig.trainer_seed. Used by "
"studentsim-reproduce table_std_seed to launch S=3 seeded "
"training runs from one base YAML."
),
)
parser.add_argument(
"--data-sampler-seed",
type=int,
default=None,
help="Override TrainingConfig.data_sampler_seed (paired with --trainer-seed).",
)
args = parser.parse_args(argv)
cfg = TrainingConfig.from_yaml(args.config)
if args.trainer_seed is not None or args.data_sampler_seed is not None:
cfg = dataclasses.replace(
cfg,
trainer_seed=(
args.trainer_seed if args.trainer_seed is not None else cfg.trainer_seed
),
data_sampler_seed=(
args.data_sampler_seed
if args.data_sampler_seed is not None
else cfg.data_sampler_seed
),
)
if cfg.stage == 1:
if args.student_id or args.roster:
parser.error("Stage-1 config; --student-id and --roster are Stage-2 only.")
trainer = Stage1Trainer(config=cfg)
return trainer.run()
# Stage 2.
if args.student_id:
per_student = _per_student_config(cfg, args.student_id)
return Stage2Trainer(config=per_student, student_id=args.student_id).run()
if args.roster is None:
parser.error("Stage-2 config; pass --student-id or --roster.")
roster = json.loads(args.roster.read_text(encoding="utf-8"))
driver = PerStudentDriver(
roster=[str(s) for s in roster],
config_builder=lambda sid: _per_student_config(cfg, sid),
)
driver.run_all(fail_fast=args.fail_fast)
if driver.failed:
return 1
return 0
if __name__ == "__main__": # pragma: no cover
sys.exit(main())