| #!/usr/bin/env python3 | |
| import json, os, subprocess, sys, torch | |
| from dataclasses import dataclass, field | |
| from typing import Optional, List | |
| import unsloth | |
| import gigatoken as gt | |
| import transformers | |
| from datasets import Dataset | |
| from peft import LoraConfig, get_peft_model | |
| from trl import SFTConfig, SFTTrainer | |
| from unsloth.chat_templates import train_on_responses_only | |
| MODEL_DIR = "/mnt/weights/scrapegoat-fp8" | |
| BUCKET = "hf://buckets/Nathan9/dump/scrapegoat-fp8" | |
| HF_CLI = os.environ.get("HF_CLI", "/home/ubuntu/.local/bin/hf") | |
| sys.path.insert(0, MODEL_DIR) | |
| from configuration_scrapegoat import ScrapeGoatConfig | |
| from modeling_scrapegoat import ScrapeGoatForCausalLM | |
| BOS_TOKEN = "[BOS]" | |
| EOS_TOKEN = "[EOS]" | |
| class ModelArguments: | |
| model_name_or_path: str = field(default=MODEL_DIR) | |
| use_lora: bool = field(default=True) | |
| lora_rank: int = field(default=16) | |
| lora_alpha: int = field(default=32) | |
| lora_dropout: float = field(default=0.05) | |
| class DataArguments: | |
| train_data_path: str = field(default=MODEL_DIR) | |
| def ensure_model_files(): | |
| needed = ["config.json", "tokenizer_config.json", "tokenization_kimi.py", | |
| "tiktoken.model", "configuration_scrapegoat.py", "modeling_scrapegoat.py", | |
| "dspark_components.py"] | |
| for f in needed: | |
| if not os.path.exists(os.path.join(MODEL_DIR, f)): | |
| src = os.path.join("/ephemeral/model_cache", f) | |
| if os.path.exists(src): | |
| os.system(f"cp {src} {os.path.join(MODEL_DIR, f)}") | |
| else: | |
| subprocess.run([HF_CLI, "cp", f"{BUCKET}/{f}", os.path.join(MODEL_DIR, f)], check=True) | |
| def ensure_shards(): | |
| TOTAL = 83 | |
| local = [f for f in os.listdir(MODEL_DIR) if f.endswith(".safetensors")] | |
| if len(local) >= TOTAL: | |
| return | |
| import glob | |
| missing = TOTAL - len(local) | |
| print(f"[shard] {missing}/{TOTAL} shards missing, starting background sync...", flush=True) | |
| subprocess.Popen(["bash", "-c", f"cd /mnt/weights && {HF_CLI} sync {BUCKET} scrapegoat-fp8/"], | |
| stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL) | |
| print(f"[shard] Sync running in background, starting model init", flush=True) | |
| def build_gigatoken_wrapper(): | |
| tk = gt.Tokenizer(MODEL_DIR).as_tiktoken() | |
| class W: | |
| def __init__(s): | |
| s._tk = tk | |
| s.eos_token = EOS_TOKEN; s.bos_token = BOS_TOKEN; s.pad_token = EOS_TOKEN | |
| s.eos_token_id = tk.encode_single_token(EOS_TOKEN) | |
| s.bos_token_id = tk.encode_single_token(BOS_TOKEN) | |
| s.pad_token_id = s.eos_token_id | |
| s.vocab_size = tk.n_vocab; s.model_max_length = 4096 | |
| def __call__(s, text, return_tensors=None, **kw): | |
| ids = s._tk.encode(text) if isinstance(text, str) else s._tk.encode_batch(text) | |
| if return_tensors == "pt": | |
| ids = torch.tensor([ids] if isinstance(text, str) else ids, dtype=torch.long) | |
| return {"input_ids": ids, "attention_mask": torch.ones_like(ids)} | |
| return {"input_ids": ids} | |
| def encode(s, text, **kw): return s._tk.encode(text) | |
| def decode(s, ids, **kw): | |
| if isinstance(ids, torch.Tensor): ids = ids.tolist() | |
| if ids and isinstance(ids[0], list): return [s._tk.decode(x) for x in ids] | |
| return s._tk.decode(ids) if isinstance(ids, list) else str(ids) | |
| def batch_decode(s, seqs, **kw): | |
| return [s._tk.decode(x.tolist() if torch.is_tensor(x) else x) for x in seqs] | |
| def apply_chat_template(s, msgs, tokenize=False, **kw): | |
| text = BOS_TOKEN | |
| for m in msgs: | |
| if m["role"] == "system": text += f"[SYSTEM] {m['content']}\n" | |
| elif m["role"] == "user": text += f"[USER] {m['content']}\n" | |
| elif m["role"] == "assistant": text += f"[ASSISTANT] {m['content']}{EOS_TOKEN}\n" | |
| return s._tk.encode(text) if tokenize else text | |
| def save_pretrained(s, path): | |
| os.makedirs(path, exist_ok=True) | |
| json.dump({"eos_token": EOS_TOKEN}, open(os.path.join(path, "tokenizer_config.json"), "w")) | |
| return W() | |
| def load_training_data(data_path, tokenizer): | |
| all_data = [] | |
| for fname in sorted(f for f in os.listdir(data_path) if f.endswith(".jsonl")): | |
| with open(os.path.join(data_path, fname)) as f: | |
| for line in f: all_data.append(json.loads(line)) | |
| def fmt(ex): | |
| return {"text": [tokenizer.apply_chat_template(m) for m in ex["messages"]]} | |
| return Dataset.from_list(all_data).map(fmt, batched=True) | |
| def load_model_zs3(ds_config_path): | |
| import deepspeed | |
| from safetensors import safe_open | |
| from transformers.integrations.deepspeed import _load_state_dict_into_zero3_model | |
| config = ScrapeGoatConfig.from_pretrained(MODEL_DIR) | |
| print("[init] Creating model under ZeRO-3 (may take a few minutes)...", flush=True) | |
| with deepspeed.zero.Init(dtype=torch.bfloat16, config_dict_or_path=ds_config_path): | |
| model = ScrapeGoatForCausalLM(config) | |
| print("[init] Model structure created, loading shards...", flush=True) | |
| TOTAL = 83 | |
| for i in range(1, TOTAL + 1): | |
| name = f"model-{i:05d}-of-{TOTAL:05d}.safetensors" | |
| p = os.path.join(MODEL_DIR, name) | |
| if not os.path.exists(p): | |
| subprocess.run([HF_CLI, "cp", f"{BUCKET}/{name}", p], check=True) | |
| sz = os.path.getsize(p) / 1e9 | |
| print(f"[load] {i}/{TOTAL}: {name} ({sz:.1f} GB)", flush=True) | |
| sd = {} | |
| with safe_open(p, framework="pt", device="cpu") as f: | |
| for k in f.keys(): sd[k] = f.get_tensor(k) | |
| _load_state_dict_into_zero3_model(model, sd, strict=False) | |
| del sd; torch.cuda.empty_cache() | |
| return model | |
| def main(): | |
| parser = transformers.HfArgumentParser((ModelArguments, DataArguments, SFTConfig)) | |
| model_args, data_args, training_args = parser.parse_args_into_dataclasses() | |
| ensure_model_files() | |
| ensure_shards() | |
| tokenizer = build_gigatoken_wrapper() | |
| script_dir = os.path.dirname(os.path.abspath(__file__)) | |
| ds_cfg = os.path.join(script_dir, "ds_config_zero3_nvme.json") | |
| model = load_model_zs3(ds_cfg) | |
| if model_args.use_lora: | |
| targets = set() | |
| for n, _ in model.named_modules(): | |
| s = n.split(".")[-1] | |
| if s in {"q_proj","k_proj","v_proj","o_proj","gate_proj","up_proj","down_proj"}: | |
| targets.add(s) | |
| model = get_peft_model(model, LoraConfig( | |
| r=model_args.lora_rank, lora_alpha=model_args.lora_alpha, | |
| lora_dropout=model_args.lora_dropout, | |
| target_modules=sorted(targets), bias="none", task_type="CAUSAL_LM", | |
| )) | |
| model.print_trainable_parameters() | |
| dataset = load_training_data(data_args.train_data_path, tokenizer) | |
| training_args.output_dir = training_args.output_dir or "/ephemeral/scrapegoat-lora" | |
| training_args.per_device_train_batch_size = 1 | |
| training_args.gradient_accumulation_steps = 8 | |
| training_args.warmup_steps = 10 | |
| training_args.max_steps = training_args.max_steps or 200 | |
| training_args.learning_rate = training_args.learning_rate or 2e-4 | |
| training_args.logging_steps = 10; training_args.save_steps = 50 | |
| training_args.optim = "adamw_torch" | |
| training_args.bf16 = True; training_args.report_to = "none" | |
| training_args.deepspeed = ds_cfg | |
| training_args.gradient_checkpointing = True | |
| training_args.gradient_checkpointing_kwargs = {"use_reentrant": False} | |
| training_args.max_seq_length = getattr(training_args, 'max_seq_length', 4096) or 4096 | |
| training_args.packing = False | |
| SFTTrainer(model=model, tokenizer=tokenizer, train_dataset=dataset, args=training_args).train() | |
| lora_dir = os.path.join(training_args.output_dir, "adapter") | |
| model.save_pretrained(lora_dir) | |
| tokenizer.save_pretrained(lora_dir) | |
| print(f"[done] LoRA saved to {lora_dir}", flush=True) | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 7.96 kB
- Xet hash:
- 47e30d2061ead02072b5bdfdcdeefa644a8d544ac3672f1f13d159a903f12632
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.