@tuned-tensor/local 0.2.3 → 0.2.5

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.
@@ -0,0 +1,270 @@
1
+ from __future__ import annotations
2
+
3
+ import inspect
4
+ import json
5
+ import os
6
+ import tarfile
7
+ import time
8
+ from pathlib import Path
9
+ from tempfile import TemporaryDirectory
10
+ from typing import Any
11
+
12
+ import torch
13
+ from datasets import Dataset as HFDataset
14
+ from peft import LoraConfig, PeftModel
15
+ from transformers import AutoModelForCausalLM, AutoTokenizer
16
+ from trl import DPOConfig, DPOTrainer
17
+
18
+
19
+ TRAINING_DIR = Path(os.environ.get("SM_CHANNEL_TRAINING", "/opt/ml/input/data/training"))
20
+ BASE_MODEL_DIR = Path(os.environ.get("SM_CHANNEL_BASE_MODEL", "/opt/ml/input/data/base_model"))
21
+ HYPERPARAMETERS_PATH = Path(
22
+ os.environ.get("TT_HYPERPARAMETERS_PATH", "/opt/ml/input/config/hyperparameters.json")
23
+ )
24
+ MODEL_DIR = Path(os.environ.get("SM_MODEL_DIR", "/opt/ml/model"))
25
+ OUTPUT_DIR = Path(os.environ.get("SM_OUTPUT_DIR", "/opt/ml/output"))
26
+
27
+
28
+ def load_hyperparameters() -> dict[str, str]:
29
+ if not HYPERPARAMETERS_PATH.is_file():
30
+ return {}
31
+ raw = json.loads(HYPERPARAMETERS_PATH.read_text())
32
+ return {str(key): str(value) for key, value in raw.items()}
33
+
34
+
35
+ HP = load_hyperparameters()
36
+
37
+
38
+ def hp(name: str, default: str | None = None) -> str | None:
39
+ value = os.getenv(f"SM_HP_{name.upper()}", HP.get(name, default))
40
+ if value is None:
41
+ return None
42
+ value = str(value).strip()
43
+ if len(value) >= 2 and value[0] == value[-1] == '"':
44
+ return value[1:-1]
45
+ return value
46
+
47
+
48
+ def hp_int(name: str, default: int) -> int:
49
+ return int(hp(name, str(default)) or default)
50
+
51
+
52
+ def hp_float(name: str, default: float) -> float:
53
+ return float(hp(name, str(default)) or default)
54
+
55
+
56
+ def hp_bool(name: str, default: bool) -> bool:
57
+ return (hp(name, str(default)) or "").lower() in {"1", "true", "yes", "y"}
58
+
59
+
60
+ def supported_kwargs(callable_obj: Any, values: dict[str, Any]) -> dict[str, Any]:
61
+ accepted = set(inspect.signature(callable_obj).parameters)
62
+ return {key: value for key, value in values.items() if key in accepted}
63
+
64
+
65
+ def load_preference_rows() -> list[dict[str, str]]:
66
+ rows: list[dict[str, str]] = []
67
+ for path in sorted(TRAINING_DIR.rglob("*.jsonl")):
68
+ for line_number, line in enumerate(path.read_text().splitlines(), start=1):
69
+ if not line.strip():
70
+ continue
71
+ row = json.loads(line)
72
+ cleaned: dict[str, str] = {}
73
+ for key in ("prompt", "chosen", "rejected"):
74
+ value = row.get(key)
75
+ if not isinstance(value, str) or not value.strip():
76
+ raise ValueError(f"{path}:{line_number} missing non-empty string field {key}")
77
+ cleaned[key] = value
78
+ rows.append(cleaned)
79
+ if not rows:
80
+ raise ValueError(f"No preference JSONL rows found under {TRAINING_DIR}")
81
+ return rows
82
+
83
+
84
+ def resolve_model_source() -> str:
85
+ archive = BASE_MODEL_DIR / "model.tar.gz"
86
+ if archive.is_file():
87
+ extracted = Path("/tmp/base-model")
88
+ extracted.mkdir(parents=True, exist_ok=True)
89
+ with tarfile.open(archive, "r:gz") as tar:
90
+ extracted_root = extracted.resolve()
91
+ for member in tar.getmembers():
92
+ destination = (extracted / member.name).resolve()
93
+ try:
94
+ destination.relative_to(extracted_root)
95
+ except ValueError:
96
+ raise ValueError(f"Unsafe archive member: {member.name}")
97
+ tar.extractall(extracted)
98
+ candidates = [path for path in extracted.iterdir() if path.is_dir()]
99
+ return str(candidates[0] if len(candidates) == 1 else extracted)
100
+ return hp("base_model") or "Qwen/Qwen3.5-2B"
101
+
102
+
103
+ def strip_file_uri(value: str | None) -> str | None:
104
+ if not value:
105
+ return value
106
+ if value.startswith("file://"):
107
+ return value[7:]
108
+ return value
109
+
110
+
111
+ def resolve_adapter_path(value: str | None, tmp: Path) -> str | None:
112
+ path_value = strip_file_uri(value)
113
+ if not path_value:
114
+ return None
115
+ path = Path(path_value)
116
+ if path.is_file() and path.name.endswith(".tar.gz"):
117
+ extracted = tmp / "parent-adapter"
118
+ extracted.mkdir(parents=True, exist_ok=True)
119
+ with tarfile.open(path, "r:gz") as tar:
120
+ extracted_root = extracted.resolve()
121
+ for member in tar.getmembers():
122
+ destination = (extracted / member.name).resolve()
123
+ try:
124
+ destination.relative_to(extracted_root)
125
+ except ValueError:
126
+ raise ValueError(f"Unsafe archive member: {member.name}")
127
+ tar.extractall(extracted)
128
+ candidates = [candidate for candidate in extracted.rglob("adapter_config.json")]
129
+ if candidates:
130
+ return str(candidates[0].parent)
131
+ directories = [candidate for candidate in extracted.iterdir() if candidate.is_dir()]
132
+ return str(directories[0] if len(directories) == 1 else extracted)
133
+ return str(path)
134
+
135
+
136
+ def create_model_and_tokenizer(model_source: str):
137
+ trust_remote_code = hp_bool("trust_remote_code", True)
138
+ token = os.getenv("HF_TOKEN")
139
+ tokenizer = AutoTokenizer.from_pretrained(
140
+ model_source,
141
+ trust_remote_code=trust_remote_code,
142
+ token=token,
143
+ )
144
+ if tokenizer.pad_token is None:
145
+ tokenizer.pad_token = tokenizer.eos_token
146
+
147
+ dtype = torch.bfloat16 if torch.cuda.is_available() and torch.cuda.is_bf16_supported() else torch.float16
148
+ model = AutoModelForCausalLM.from_pretrained(
149
+ model_source,
150
+ trust_remote_code=trust_remote_code,
151
+ token=token,
152
+ torch_dtype=dtype if torch.cuda.is_available() else None,
153
+ device_map="auto" if torch.cuda.is_available() else None,
154
+ )
155
+ model.config.use_cache = False
156
+ return model, tokenizer
157
+
158
+
159
+ def lora_target_modules(default: str) -> list[str] | str:
160
+ raw = hp("lora_target_modules", default) or default
161
+ if raw == "all-linear":
162
+ return raw
163
+ return [item.strip() for item in raw.split(",") if item.strip()]
164
+
165
+
166
+ def lora_config() -> LoraConfig:
167
+ return LoraConfig(
168
+ r=hp_int("lora_rank", 16),
169
+ lora_alpha=hp_int("lora_alpha", 32),
170
+ lora_dropout=hp_float("lora_dropout", 0.05),
171
+ bias="none",
172
+ task_type="CAUSAL_LM",
173
+ target_modules=lora_target_modules("all-linear"),
174
+ )
175
+
176
+
177
+ def create_model_archive() -> Path:
178
+ OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
179
+ archive_path = OUTPUT_DIR / "model.tar.gz"
180
+ with tarfile.open(archive_path, "w:gz") as tar:
181
+ tar.add(MODEL_DIR, arcname="model")
182
+ return archive_path
183
+
184
+
185
+ def dpo_config(output_dir: str) -> DPOConfig:
186
+ max_prompt_length = hp("max_prompt_length")
187
+ max_completion_length = hp("max_completion_length")
188
+ config_values: dict[str, Any] = {
189
+ "output_dir": output_dir,
190
+ "num_train_epochs": hp_int("n_epochs", 3),
191
+ "learning_rate": hp_float("learning_rate", 0.00001),
192
+ "per_device_train_batch_size": hp_int("per_device_train_batch_size", 1),
193
+ "gradient_accumulation_steps": hp_int("gradient_accumulation_steps", 8),
194
+ "logging_steps": 1,
195
+ "save_strategy": "no",
196
+ "report_to": "none",
197
+ "remove_unused_columns": False,
198
+ "bf16": torch.cuda.is_available() and torch.cuda.is_bf16_supported(),
199
+ "fp16": torch.cuda.is_available() and not torch.cuda.is_bf16_supported(),
200
+ "gradient_checkpointing": True,
201
+ "gradient_checkpointing_kwargs": {"use_reentrant": False},
202
+ "beta": hp_float("dpo_beta", 0.1),
203
+ "loss_type": hp("dpo_loss_type", "sigmoid"),
204
+ "label_smoothing": hp_float("dpo_label_smoothing", 0.0),
205
+ "reference_free": hp_bool("dpo_reference_free", False),
206
+ "max_length": hp_int("max_seq_length", 2048),
207
+ "max_prompt_length": int(max_prompt_length) if max_prompt_length else None,
208
+ "max_completion_length": int(max_completion_length) if max_completion_length else None,
209
+ }
210
+ return DPOConfig(**supported_kwargs(
211
+ DPOConfig,
212
+ {key: value for key, value in config_values.items() if value is not None},
213
+ ))
214
+
215
+
216
+ def run_training(rows: list[dict[str, str]], model_source: str) -> dict[str, Any]:
217
+ with TemporaryDirectory() as tmp:
218
+ parent_adapter = resolve_adapter_path(hp("parent_model_artifact"), Path(tmp))
219
+ model, tokenizer = create_model_and_tokenizer(model_source)
220
+ if parent_adapter:
221
+ model = PeftModel.from_pretrained(model, parent_adapter, is_trainable=True)
222
+ dataset = HFDataset.from_list(rows)
223
+
224
+ args = dpo_config(tmp)
225
+ trainer_values: dict[str, Any] = {
226
+ "model": model,
227
+ "ref_model": None,
228
+ "args": args,
229
+ "train_dataset": dataset,
230
+ "processing_class": tokenizer,
231
+ "tokenizer": tokenizer,
232
+ }
233
+ if not parent_adapter:
234
+ trainer_values["peft_config"] = lora_config()
235
+ trainer = DPOTrainer(**supported_kwargs(DPOTrainer, trainer_values))
236
+ result = trainer.train()
237
+ trainer.save_model(MODEL_DIR)
238
+
239
+ tokenizer.save_pretrained(MODEL_DIR)
240
+ return result.metrics
241
+
242
+
243
+ def main() -> None:
244
+ started = time.time()
245
+ MODEL_DIR.mkdir(parents=True, exist_ok=True)
246
+ OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
247
+
248
+ rows = load_preference_rows()
249
+ model_source = resolve_model_source()
250
+ metrics = run_training(rows, model_source)
251
+
252
+ archive_path = create_model_archive()
253
+ output_metrics = {
254
+ "training_method": "dpo",
255
+ "preference_rows": len(rows),
256
+ "model_source": model_source,
257
+ "parent_model_artifact": hp("parent_model_artifact"),
258
+ "dpo_beta": hp_float("dpo_beta", 0.1),
259
+ "dpo_loss_type": hp("dpo_loss_type", "sigmoid"),
260
+ "train_runtime": round(time.time() - started, 3),
261
+ **{key: float(value) for key, value in metrics.items() if isinstance(value, (int, float))},
262
+ "model_archive": str(archive_path),
263
+ }
264
+ (MODEL_DIR / "training-metrics.json").write_text(json.dumps(output_metrics, indent=2))
265
+ (OUTPUT_DIR / "training-metrics.json").write_text(json.dumps(output_metrics, indent=2))
266
+ print(json.dumps(output_metrics, indent=2))
267
+
268
+
269
+ if __name__ == "__main__":
270
+ main()