| |
| """ |
| Phase 3 — QLoRA fine-tune of Muse-Glimmer-30B (text-only, perception encoder frozen). |
| |
| Run entirely inside WSL, from the venv at ~/glimmer/venv. See scripts/TRAINING.md |
| for exact launch / pause / resume / monitor commands. |
| |
| Requirements this implements (see briefs/task-3-brief.md for the authoritative spec): |
| 1. Load via Unsloth (FastLanguageModel, falling back to FastVisionModel) from the |
| local HF cache only -- never re-downloads, asserts a cache hit up front. |
| 2. Freezes the perception encoder (vision tower + vision->text projector): prefers |
| Unsloth's own finetune_vision_layers=False-style flag, then independently |
| verifies via named_parameters() that zero trainable params carry a vision/ |
| projector-ish name. |
| 3. LoRA r=16/alpha=16/dropout=0 on the text tower's attention + MLP projections |
| only (q/k/v/o, gate/up/down) -- no embeddings, no lm_head. |
| 4. Formats dataset/dataset.jsonl with tokenizer.apply_chat_template (bespoke |
| <|start|>/<|message|>/<|eot|> template, reasoning_strength='high'), never |
| hand-rolled. |
| 5. Masks non-assistant turns via unsloth_zoo's train_on_responses_only, using the |
| template's real turn markers (derived from chat_template.jinja + confirmed |
| against tokenizer_config.json's response_template, and verified for real |
| against 2 decoded dataset samples -- see briefs/task-3-report.md). |
| 6. Trainer config is plan-mandated (see brief); checkpointing is tightened to |
| save_steps=200 / save_total_limit=3 for pausability. |
| 7-9. Pausable training: a TrainerCallback polls ~/glimmer/PAUSE on_step_end, saves |
| a checkpoint, logs "PAUSED at step N", stops cleanly, and deletes the sentinel |
| itself. --resume picks up the latest checkpoint (weights + optimizer + scheduler |
| + step). |
| 10. GPU strategy / OOM ladder is a launch-time concern (CUDA_VISIBLE_DEVICES / batch |
| size / device_map), documented in TRAINING.md -- this script exposes the knobs |
| as CLI flags rather than hardcoding one config. |
| 11-13. Smoke leg: actually run at --max_steps 20 --max_seq_length 2048 (rung (c) -- |
| rung (a)'s seq 4096 measured ~2550s/step and was abandoned as impractical; |
| 100 steps was the original target but a long resumed run was killed by |
| something external to this script partway through -- see the report for the |
| full, honest chain of real-world adaptations). Loss series, one real |
| pause/resume cycle with verified step continuity, 3 held-out generations, |
| and a timing projection are all in briefs/task-3-report.md with real numbers. |
| |
| RESOLVED BLOCKER (was open when this script was first written -- see |
| briefs/task-3-report.md for full history): `transformers==5.5.0` did not register the |
| `muse_glimmer` model_type in CONFIG_MAPPING_NAMES at all. The orchestrator upgraded |
| the shared venv to `transformers==5.15.0`, which does register it; `FastLanguageModel. |
| from_pretrained` now loads the model (confirmed live, 22.2GB weights-only on one |
| 3090). One live gotcha this upgrade surfaced: the tokenizer object this model's |
| `from_pretrained` returns is a MuseGlimmerProcessor, not a plain tokenizer -- see |
| `get_text_tokenizer()`. |
| """ |
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import os |
| import sys |
| import time |
| from pathlib import Path |
|
|
| |
| |
| |
| |
| |
| GLIMMER_HOME = Path(os.environ.get("GLIMMER_HOME", os.path.expanduser("~/glimmer"))) |
| os.environ.setdefault("HF_HOME", str(GLIMMER_HOME / "hf_home")) |
| os.environ["HF_HUB_OFFLINE"] = "1" |
| os.environ["TOKENIZERS_PARALLELISM"] = "false" |
|
|
| MODEL_NAME = "meta-models/Muse-Glimmer-30B" |
| DATASET_PATH_DEFAULT = GLIMMER_HOME / "dataset" / "dataset.jsonl" |
| OUTPUT_DIR_DEFAULT = GLIMMER_HOME / "runs" / "sentry-v01" |
| PAUSE_SENTINEL_DEFAULT = GLIMMER_HOME / "PAUSE" |
| HOLDOUT_DIR_DEFAULT = Path("/mnt/c/Users/Dwain-Admin/Desktop/GLIMMER-SENTRY-30B/dataset/holdout") |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| INSTRUCTION_PART = "<|start|>user<|message|>" |
| RESPONSE_PART = "<|start|>assistant to=user<|message|>" |
|
|
| REASONING_STRENGTH = "high" |
|
|
| |
| |
| |
| |
| |
| |
| VISION_NAME_FRAGMENTS = [ |
| "vision", "visual", "projector", "vision_tower", "multi_modal_projector", |
| "image_newline", "patch_embed", "vit.", ".vit", "perceiver", |
| ] |
|
|
| LORA_R = 16 |
| LORA_ALPHA = 16 |
| LORA_DROPOUT = 0.0 |
| |
| |
| |
| |
| TEXT_TOWER_TARGET_MODULES = [ |
| "q_proj", "k_proj", "v_proj", "o_proj", |
| "gate_proj", "up_proj", "down_proj", |
| ] |
|
|
|
|
| def log(msg: str) -> None: |
| print(f"[train.py {time.strftime('%H:%M:%S')}] {msg}", flush=True) |
|
|
|
|
| def get_text_tokenizer(tokenizer_or_processor): |
| """FastLanguageModel.from_pretrained returns a MuseGlimmerProcessor for this |
| model (confirmed live, transformers 5.15.0), not a plain tokenizer -- its |
| __call__'s first positional arg is `images`, so calling it positionally with |
| text (`tokenizer(text)`) raises a base64/"Incorrect padding" error deep in |
| the image-processing path instead of tokenizing. The plain tokenizer used for |
| every text-only operation here (apply_chat_template, train_on_responses_only, |
| the trainer's processing_class, generation) is the inner `.tokenizer` |
| attribute. Falls back to the object itself if it's already a plain tokenizer |
| (no `.tokenizer` attribute) so this is safe either way.""" |
| return getattr(tokenizer_or_processor, "tokenizer", tokenizer_or_processor) |
|
|
|
|
| |
| |
| |
| def assert_cache_hit(model_name: str) -> str: |
| """Resolve the model snapshot from the local HF cache only. Raises loudly |
| (via HF_HUB_OFFLINE=1, already set above) instead of downloading anything.""" |
| from huggingface_hub import snapshot_download |
|
|
| try: |
| path = snapshot_download(model_name, local_files_only=True) |
| except Exception as e: |
| raise RuntimeError( |
| f"Cache-hit assertion FAILED for '{model_name}' under HF_HOME=" |
| f"{os.environ['HF_HOME']}. Refusing to download (HF_HUB_OFFLINE=1). " |
| f"Original error: {e}" |
| ) from e |
| log(f"Cache hit confirmed: {model_name} -> {path}") |
| return path |
|
|
|
|
| |
| |
| |
| def load_model_and_tokenizer(max_seq_length: int, device_map: str): |
| """Try FastLanguageModel first, fall back to FastVisionModel, per the brief. |
| Returns (model, tokenizer, loader_name). Raises with the exact upstream error |
| (never falls back to raw transformers+peft) if both fail -- that is a BLOCKED |
| condition per the brief, not something to improvise around.""" |
| import torch |
| from unsloth import FastLanguageModel, FastVisionModel |
|
|
| common_kwargs = dict( |
| model_name=MODEL_NAME, |
| max_seq_length=max_seq_length, |
| load_in_4bit=True, |
| dtype=torch.bfloat16, |
| device_map=device_map, |
| ) |
|
|
| log("Attempting FastLanguageModel.from_pretrained ...") |
| try: |
| model, tokenizer = FastLanguageModel.from_pretrained(**common_kwargs) |
| return model, tokenizer, "FastLanguageModel" |
| except Exception as e_lang: |
| log(f"FastLanguageModel failed: {e_lang!r}") |
| log("Falling back to FastVisionModel.from_pretrained ...") |
| try: |
| model, tokenizer = FastVisionModel.from_pretrained(**common_kwargs) |
| return model, tokenizer, "FastVisionModel" |
| except Exception as e_vision: |
| raise RuntimeError( |
| "BLOCKED: Muse-Glimmer-30B could not be loaded by either " |
| "FastLanguageModel or FastVisionModel in unsloth 2026.8.12 / " |
| f"transformers {__import__('transformers').__version__}.\n\n" |
| f"FastLanguageModel error:\n{e_lang!r}\n\n" |
| f"FastVisionModel error:\n{e_vision!r}\n\n" |
| "Per the brief: do not improvise with raw transformers+peft here -- " |
| "this is a BLOCKED condition to report, not to route around." |
| ) from e_vision |
|
|
|
|
| def apply_lora_with_frozen_vision(model, loader_name: str, use_gradient_checkpointing="unsloth"): |
| """Attach LoRA to the text tower only. Prefers Unsloth's own |
| finetune_vision_layers=False-style flag (FastVisionModel.get_peft_model) -- |
| tried regardless of which loader actually succeeded, since that flag operates |
| on the already-loaded model object via get_peft_regex's own module |
| introspection, not on loader-stamped state. Falls back to the explicit |
| text-tower target_modules list (FastLanguageModel.get_peft_model / plain |
| LoraConfig) if the flag-based call raises (e.g. genuine incompatibility with |
| a FastLanguageModel-loaded object). named_parameters() is independently |
| verified afterward regardless of which path was used -- see verify_freeze().""" |
| from unsloth import FastLanguageModel, FastVisionModel |
|
|
| try: |
| log("Attempting FastVisionModel.get_peft_model with finetune_vision_layers=False " |
| "(Unsloth's native vision-freeze flag) ...") |
| model = FastVisionModel.get_peft_model( |
| model, |
| r=LORA_R, |
| lora_alpha=LORA_ALPHA, |
| lora_dropout=LORA_DROPOUT, |
| bias="none", |
| finetune_vision_layers=False, |
| finetune_language_layers=True, |
| finetune_attention_modules=True, |
| finetune_mlp_modules=True, |
| use_gradient_checkpointing=use_gradient_checkpointing, |
| random_state=42, |
| ) |
| log("FastVisionModel.get_peft_model succeeded.") |
| except Exception as e: |
| log(f"FastVisionModel.get_peft_model failed ({e!r}); falling back to the " |
| f"explicit text-tower target_modules list via FastLanguageModel.get_peft_model.") |
| model = FastLanguageModel.get_peft_model( |
| model, |
| r=LORA_R, |
| target_modules=TEXT_TOWER_TARGET_MODULES, |
| lora_alpha=LORA_ALPHA, |
| lora_dropout=LORA_DROPOUT, |
| bias="none", |
| use_gradient_checkpointing=use_gradient_checkpointing, |
| random_state=42, |
| ) |
| return model |
|
|
|
|
| def verify_freeze(model) -> dict: |
| """Independent, from-scratch verification (not trusting the flag above): |
| walk named_parameters(), tally total/trainable, and confirm zero trainable |
| params carry a vision/projector-ish name. Defensively sets requires_grad=False |
| on any that slip through, and prints + returns everything for the report.""" |
| total_params = 0 |
| trainable_params = 0 |
| trainable_vision_params = 0 |
| matched_prefixes = set() |
| offending_names = [] |
|
|
| for name, p in model.named_parameters(): |
| n = p.numel() |
| total_params += n |
| lname = name.lower() |
| is_vision_ish = any(frag in lname for frag in VISION_NAME_FRAGMENTS) |
| if p.requires_grad: |
| trainable_params += n |
| if is_vision_ish: |
| trainable_vision_params += n |
| offending_names.append(name) |
| matched_prefixes.add(".".join(name.split(".")[:4])) |
| |
| p.requires_grad_(False) |
| trainable_params -= n |
|
|
| result = { |
| "total_params": total_params, |
| "trainable_params": trainable_params, |
| "trainable_vision_params_before_defensive_freeze": trainable_vision_params, |
| "offending_names_sample": offending_names[:20], |
| } |
| log(f"Freeze verification: total_params={total_params:,} " |
| f"trainable_params={trainable_params:,} " |
| f"trainable_pct={100*trainable_params/max(total_params,1):.4f}%") |
| if trainable_vision_params > 0: |
| log(f"WARNING: {trainable_vision_params:,} trainable params matched a " |
| f"vision/projector name pattern and were forcibly frozen just now. " |
| f"Prefixes: {sorted(matched_prefixes)}") |
| else: |
| log("Confirmed: zero trainable params match a vision/projector name pattern.") |
| return result |
|
|
|
|
| |
| |
| |
| def load_and_format_dataset(dataset_path: Path, tokenizer): |
| from datasets import load_dataset |
|
|
| log(f"Loading dataset from {dataset_path}") |
| ds = load_dataset("json", data_files=str(dataset_path), split="train") |
| log(f"Loaded {len(ds)} examples") |
|
|
| def _format(example): |
| text = tokenizer.apply_chat_template( |
| example["messages"], |
| tokenize=False, |
| add_generation_prompt=False, |
| reasoning_strength=REASONING_STRENGTH, |
| ) |
| return {"text": text} |
|
|
| ds = ds.map(_format, remove_columns=[c for c in ds.column_names if c != "messages"]) |
| log("Formatted dataset with tokenizer.apply_chat_template " |
| f"(reasoning_strength='{REASONING_STRENGTH}')") |
| log(f"Sample formatted example (first 800 chars):\n{ds[0]['text'][:800]}") |
| return ds |
|
|
|
|
| |
| |
| |
| def apply_response_masking(trainer): |
| """Wrap trainer.train_dataset so only assistant turns contribute to the loss, |
| using the template's real turn markers. See briefs/task-3-report.md for the |
| standalone (tokenizer-only) verification that this masks correctly across a |
| 2-turn and a multi-turn example before this was ever wired into a trainer.""" |
| from unsloth_zoo.dataset_utils import train_on_responses_only |
|
|
| log(f"Applying train_on_responses_only: instruction_part={INSTRUCTION_PART!r} " |
| f"response_part={RESPONSE_PART!r}") |
| trainer = train_on_responses_only( |
| trainer, |
| instruction_part=INSTRUCTION_PART, |
| response_part=RESPONSE_PART, |
| ) |
| return trainer |
|
|
|
|
| |
| |
| |
| def make_pause_callback(sentinel_path: Path): |
| """Builds a TrainerCallback subclass at call time (transformers is imported |
| lazily, matching the rest of this script's import style) that polls the |
| PAUSE sentinel file on_step_end. |
| |
| CAUGHT LIVE DURING THE SMOKE LEG (see report): an earlier version of this |
| defined `class _PauseCallback(TrainerCallback, PauseCallback)` -- a plain |
| mixin combined via multiple inheritance -- which is a real Python MRO trap: |
| `TrainerCallback` (listed first) defines its OWN no-op `on_step_end` stub, |
| which method resolution order finds before the real implementation further |
| down the MRO, silently shadowing it. The callback was then a structural |
| no-op: `trainer.add_callback(...)` succeeded, training ran, but on_step_end |
| never actually executed the sentinel check, so touching PAUSE mid-run had |
| no effect at all -- confirmed live (sentinel sat untouched for 2+ full |
| steps past creation). Fixed by inheriting from TrainerCallback directly, |
| which is also just simpler.""" |
| from transformers import TrainerCallback |
|
|
| class PauseCallback(TrainerCallback): |
| """on_step_end: if the sentinel exists, force an immediate checkpoint |
| save, log a clear PAUSED line, stop training cleanly, and delete the |
| sentinel itself (so the next launch doesn't immediately re-pause).""" |
|
|
| def __init__(self, sentinel_path: Path): |
| self.sentinel_path = Path(sentinel_path) |
|
|
| def on_step_end(self, args, state, control, **kwargs): |
| if self.sentinel_path.exists(): |
| log(f"PAUSE sentinel found at {self.sentinel_path} -- pausing at " |
| f"step {state.global_step}.") |
| control.should_save = True |
| control.should_training_stop = True |
| try: |
| self.sentinel_path.unlink() |
| log(f"Deleted PAUSE sentinel {self.sentinel_path}") |
| except FileNotFoundError: |
| pass |
| log(f"PAUSED at step {state.global_step}") |
| return control |
|
|
| return PauseCallback(sentinel_path) |
|
|
|
|
| |
| |
| |
| def build_holdout_prompts(holdout_dir: Path) -> list[dict]: |
| """Builds exactly the 3 required held-out prompts (translation Sigma->KQL, |
| explanation, authoring), all sourced from dataset/holdout/*.yml so none of |
| them can have leaked into training. Prompt phrasing mirrors |
| scripts/build_dataset.py's own template banks so the smoke-test prompts are |
| representative of the training distribution.""" |
| yml_files = sorted(holdout_dir.glob("*.yml")) |
| if len(yml_files) < 3: |
| raise RuntimeError(f"Expected >=3 holdout files in {holdout_dir}, found {len(yml_files)}") |
|
|
| import yaml |
|
|
| def _load(path): |
| raw = path.read_text(encoding="utf-8") |
| parsed = yaml.safe_load(raw) |
| return raw, parsed |
|
|
| raw0, rule0 = _load(yml_files[0]) |
| raw1, rule1 = _load(yml_files[1]) |
| raw2, rule2 = _load(yml_files[2]) |
|
|
| translation_prompt = ( |
| f"Convert this Sigma rule to Microsoft 365 Defender Advanced Hunting KQL:\n\n" |
| f"```yaml\n{raw0}\n```" |
| ) |
| explanation_prompt = f"Explain this Sigma rule in plain English:\n\n```yaml\n{raw1}\n```" |
| logsource_str = ", ".join(f"{k}={v}" for k, v in (rule2.get("logsource") or {}).items()) |
| tags_str = ", ".join(rule2.get("tags") or []) |
| authoring_prompt = ( |
| f"Write a Sigma rule that detects: {rule2.get('description', rule2.get('title', ''))}\n\n" |
| f"Logsource: {logsource_str}\nRelevant ATT&CK tags: {tags_str}" |
| ) |
|
|
| return [ |
| {"task": "translation_sigma_to_kql", "source_file": yml_files[0].name, "prompt": translation_prompt}, |
| {"task": "explanation", "source_file": yml_files[1].name, "prompt": explanation_prompt}, |
| {"task": "authoring", "source_file": yml_files[2].name, "prompt": authoring_prompt}, |
| ] |
|
|
|
|
| def run_holdout_generations(model, tokenizer, holdout_dir: Path, out_path: Path): |
| from unsloth import FastLanguageModel |
|
|
| prompts = build_holdout_prompts(holdout_dir) |
| FastLanguageModel.for_inference(model) |
|
|
| results = [] |
| for item in prompts: |
| messages = [{"role": "user", "content": item["prompt"]}] |
| inputs = tokenizer.apply_chat_template( |
| messages, |
| tokenize=True, |
| add_generation_prompt=True, |
| reasoning_strength=REASONING_STRENGTH, |
| return_tensors="pt", |
| ).to(model.device) |
| t0 = time.time() |
| out_ids = model.generate( |
| input_ids=inputs, |
| max_new_tokens=512, |
| do_sample=True, |
| temperature=1.0, |
| top_p=0.95, |
| top_k=64, |
| ) |
| gen_text = tokenizer.decode(out_ids[0][inputs.shape[1]:], skip_special_tokens=False) |
| elapsed = time.time() - t0 |
| log(f"Generated for task={item['task']} in {elapsed:.1f}s") |
| results.append({**item, "generation": gen_text, "seconds": elapsed}) |
|
|
| out_path.write_text(json.dumps(results, indent=2), encoding="utf-8") |
| log(f"Wrote {len(results)} held-out generations to {out_path}") |
| return results |
|
|
|
|
| |
| |
| |
| def main(): |
| ap = argparse.ArgumentParser(description=__doc__) |
| ap.add_argument("--dataset_path", type=Path, default=DATASET_PATH_DEFAULT) |
| ap.add_argument("--output_dir", type=Path, default=OUTPUT_DIR_DEFAULT) |
| ap.add_argument("--holdout_dir", type=Path, default=HOLDOUT_DIR_DEFAULT) |
| ap.add_argument("--pause_sentinel", type=Path, default=PAUSE_SENTINEL_DEFAULT) |
| ap.add_argument("--max_seq_length", type=int, default=4096) |
| |
| |
| |
| |
| |
| |
| |
| ap.add_argument("--per_device_train_batch_size", type=int, default=1) |
| ap.add_argument("--gradient_accumulation_steps", type=int, default=16) |
| ap.add_argument("--device_map", type=str, default="sequential", |
| help="'sequential' for single-GPU (use with CUDA_VISIBLE_DEVICES=1 -- " |
| "GPU 0 carries ~1-2GB of desktop apps, GPU 1 is clean), " |
| "'balanced' to shard across both GPUs (OOM fallback rung (b)).") |
| ap.add_argument("--max_steps", type=int, default=None, |
| help="The smoke leg actually used 20 (see briefs/task-3-report.md " |
| "for why, not the originally-planned 100) at " |
| "--max_seq_length 2048 (~229 steps/epoch at that config). " |
| "Omit for a full run (num_train_epochs=1) -- NOT invoked by " |
| "this task.") |
| ap.add_argument("--resume", action="store_true", |
| help="Resume from the latest checkpoint in --output_dir " |
| "(weights + optimizer + scheduler + step).") |
| ap.add_argument("--skip_generation", action="store_true", |
| help="Skip the post-training held-out generation step.") |
| ap.add_argument("--save_steps", type=int, default=200, |
| help="Checkpoint interval in optimizer steps. The full run " |
| "uses a tight interval (e.g. 20) so an external kill " |
| "costs at most ~40 min at rung (c) step times.") |
| ap.add_argument("--warmup_steps", type=int, default=None, |
| help="Explicit override for warmup step count, used only on " |
| "installs where SFTConfig has no warmup_ratio field (see " |
| "the warmup_ratio/warmup_steps translation note below). " |
| "Recommended for a full run once the real packed dataset " |
| "size is known; the smoke leg derives it from --max_steps.") |
| args = ap.parse_args() |
|
|
| log(f"HF_HOME={os.environ['HF_HOME']} HF_HUB_OFFLINE={os.environ['HF_HUB_OFFLINE']}") |
| assert_cache_hit(MODEL_NAME) |
|
|
| model, tokenizer, loader_name = load_model_and_tokenizer( |
| max_seq_length=args.max_seq_length, device_map=args.device_map |
| ) |
| log(f"Loaded via {loader_name} (raw processing object type: {type(tokenizer).__name__})") |
|
|
| |
| |
| |
| |
| |
| tokenizer = get_text_tokenizer(tokenizer) |
| log(f"Using plain text tokenizer for all downstream ops: {type(tokenizer).__name__}") |
|
|
| model = apply_lora_with_frozen_vision(model, loader_name) |
| freeze_stats = verify_freeze(model) |
|
|
| ds = load_and_format_dataset(args.dataset_path, tokenizer) |
|
|
| from trl import SFTConfig, SFTTrainer |
|
|
| args.output_dir.mkdir(parents=True, exist_ok=True) |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| import inspect as _inspect |
| STEPS_PER_EPOCH_ESTIMATE = 115 |
| warmup_ratio = 0.03 |
| supports_warmup_ratio = "warmup_ratio" in _inspect.signature(SFTConfig.__init__).parameters |
| warmup_kwarg = {} |
| if supports_warmup_ratio: |
| warmup_kwarg["warmup_ratio"] = warmup_ratio |
| log("warmup_ratio is supported natively by the installed TRL/transformers.") |
| elif args.warmup_steps is not None: |
| warmup_kwarg["warmup_steps"] = args.warmup_steps |
| log(f"Using explicit --warmup_steps={args.warmup_steps} (warmup_ratio unsupported).") |
| else: |
| total_steps_for_warmup = args.max_steps if args.max_steps else STEPS_PER_EPOCH_ESTIMATE |
| warmup_steps = max(1, round(warmup_ratio * total_steps_for_warmup)) |
| warmup_kwarg["warmup_steps"] = warmup_steps |
| import transformers as _tf |
| log(f"WARNING: installed TrainingArguments (transformers {_tf.__version__}) has no " |
| f"warmup_ratio field -- translating the brief's warmup_ratio=0.03 into " |
| f"warmup_steps={warmup_steps} (3% of {total_steps_for_warmup} total steps" |
| + ("" if args.max_steps else |
| " [steps/epoch ESTIMATE from token-count analysis in the report -- pass " |
| "--warmup_steps explicitly for a precise full-run value once the real packed " |
| "dataset size is known]") + ").") |
|
|
| sft_config = SFTConfig( |
| output_dir=str(args.output_dir), |
| per_device_train_batch_size=args.per_device_train_batch_size, |
| gradient_accumulation_steps=args.gradient_accumulation_steps, |
| learning_rate=2e-4, |
| num_train_epochs=1, |
| max_steps=args.max_steps if args.max_steps else -1, |
| optim="adamw_8bit", |
| bf16=True, |
| gradient_checkpointing=True, |
| lr_scheduler_type="cosine", |
| **warmup_kwarg, |
| seed=42, |
| logging_steps=5, |
| save_steps=args.save_steps, |
| save_total_limit=3, |
| packing=True, |
| max_length=args.max_seq_length, |
| dataset_text_field="text", |
| report_to="none", |
| ) |
|
|
| trainer = SFTTrainer( |
| model=model, |
| processing_class=tokenizer, |
| train_dataset=ds, |
| args=sft_config, |
| ) |
| trainer = apply_response_masking(trainer) |
| trainer.add_callback(make_pause_callback(args.pause_sentinel)) |
|
|
| import torch |
| torch.cuda.reset_peak_memory_stats() |
| t_train_start = time.time() |
| trainer.train(resume_from_checkpoint=True if args.resume else False) |
| train_wall_s = time.time() - t_train_start |
| peak_vram = torch.cuda.max_memory_allocated() / (1024 ** 3) |
| log(f"Training loop finished/paused. Wall clock: {train_wall_s:.1f}s. " |
| f"Peak VRAM: {peak_vram:.2f} GiB") |
|
|
| if not args.skip_generation: |
| run_holdout_generations( |
| model, tokenizer, args.holdout_dir, |
| args.output_dir / "holdout_generations.json", |
| ) |
|
|
| log("Done.") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|