Nathan9/dump / scrapegoat-lora /train_unsloth.py
Nathan9's picture
download
raw
7.96 kB
#!/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]"
@dataclass
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)
@dataclass
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.