import os import random import json import shutil from dataclasses import dataclass from typing import Any, Dict, List, Tuple from datasets import load_dataset, Dataset from transformers import TrainingArguments, Trainer, TrainerCallback import numpy as np import torch from unsloth import FastLanguageModel def _getenv(name, default): return os.environ.get(name, default) def _getenv_int(name, default): try: return int(os.environ.get(name, str(default))) except: return default def _getenv_float(name, default): try: return float(os.environ.get(name, str(default))) except: return default BASE_MODEL_ID = _getenv("SFT_BASE_MODEL", "Qwen/Qwen3-4B-Instruct-2507") DATASET_ID = _getenv("SFT_DATASET_ID", "u-10bei/structured_data_with_cot_dataset_512_v4") OUT_LORA_DIR = _getenv("SFT_OUT_LORA_DIR", "/kaggle/working/lora_structeval_t_qwen3_4b") SEED = _getenv_int("SFT_SEED", 3407) VAL_RATIO = _getenv_float("SFT_VAL_RATIO", 0.05) MAX_SEQ_LEN = _getenv_int("SFT_MAX_SEQ_LEN", 512) LORA_R = _getenv_int("SFT_LORA_R", 64) LORA_ALPHA = _getenv_int("SFT_LORA_ALPHA", 128) LORA_DROPOUT = _getenv_float("SFT_LORA_DROPOUT", 0) LORA_TARGET_MODULES = _getenv("SFT_LORA_TARGET_MODULES", "q_proj,k_proj,v_proj,o_proj,gate_proj,up_proj,down_proj").split(",") NUM_TRAIN_EPOCHS = _getenv_int("SFT_EPOCHS", 1) PER_DEVICE_TRAIN_BATCH_SIZE = _getenv_int("SFT_PER_DEVICE_TRAIN_BS", 2) PER_DEVICE_EVAL_BATCH_SIZE = _getenv_int("SFT_PER_DEVICE_EVAL_BS", 2) GRAD_ACCUM = _getenv_int("SFT_GRAD_ACCUM", 8) LR = _getenv_float("SFT_LR", 1e-6) WARMUP_RATIO = _getenv_float("SFT_WARMUP_RATIO", 0.1) MAX_STEPS = _getenv_int("SFT_MAX_STEPS", -1) LOGGING_STEPS = _getenv_int("SFT_LOGGING_STEPS", 10) EVAL_STEPS = _getenv_int("SFT_EVAL_STEPS", 50) SAVE_STEPS = _getenv_int("SFT_SAVE_STEPS", 100) SAVE_TOTAL_LIMIT = _getenv_int("SFT_SAVE_TOTAL_LIMIT", 2) WEIGHT_DECAY = _getenv_float("SFT_WEIGHT_DECAY", 0.05) UPSAMPLE_ENABLE = _getenv("SFT_USE_UPSAMPLING", "0") in ("1","true","True") UPSAMPLE_RULES_JSON = _getenv("SFT_UPSAMPLE_RULES", "") def seed_everything(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) seed_everything(SEED) def ensure_openai_messages(ds, msg_col="messages"): ex = ds[0].get(msg_col, None) if not isinstance(ex, list): raise ValueError(f"Dataset must have list-style messages. Got {type(ex)}") def has_any_nonempty_assistant_turn(msgs): return any(m.get("role")=="assistant" and str(m.get("content","")).strip()!="" for m in msgs) def ends_with_nonempty_assistant(ex): msgs = ex.get("messages", []) if not msgs or msgs[-1].get("role")!="assistant": return False c = msgs[-1].get("content","") return isinstance(c, str) and c.strip()!="" def shuffle_split(ds, val_ratio, seed): ds_shuf = ds.shuffle(seed=seed) n = len(ds_shuf) n_val = max(1, int(round(n * val_ratio))) return ds_shuf.select(range(n_val, n)), ds_shuf.select(range(n_val)) def make_text_cache_builder(tokenizer): def _build(batch): full_out, prefix_out, full_len_out, prefix_len_out = [], [], [], [] for msgs in batch["messages"]: full = tokenizer.apply_chat_template(msgs, tokenize=False, add_generation_prompt=False) prefix = tokenizer.apply_chat_template(msgs[:-1], tokenize=False, add_generation_prompt=True) full_out.append(full); prefix_out.append(prefix) full_ids = tokenizer(full, add_special_tokens=False, truncation=False)["input_ids"] prefix_ids = tokenizer(prefix, add_special_tokens=False, truncation=False)["input_ids"] full_len_out.append(len(full_ids)); prefix_len_out.append(len(prefix_ids)) return {"full_text": full_out, "prefix_text": prefix_out, "full_input_ids_len": full_len_out, "prefix_input_ids_len": prefix_len_out} return _build MASK_COT = _getenv("SFT_MASK_COT", "1") in ("1","true","True") OUTPUT_MARKERS = [s.strip() for s in _getenv("SFT_OUTPUT_MARKERS", "Output:,OUTPUT:,Final:,Answer:,Result:,Response:").split(",") if s.strip()] OUTPUT_LEARN_MODE = _getenv("SFT_OUTPUT_LEARN_MODE", "after_marker") @dataclass class AssistantOnlyCollatorCached: tokenizer: Any max_length: int = MAX_SEQ_LEN def _find_subseq(self, seq, sub): if not sub or len(sub) > len(seq): return -1 for i in range(len(seq) - len(sub) + 1): if seq[i:i+len(sub)] == sub: return i return -1 def __call__(self, batch): tok = self.tokenizer full_texts = [ex["full_text"] for ex in batch] prefix_texts = [ex["prefix_text"] for ex in batch] old_trunc = getattr(tok, "truncation_side", "right") old_pad = getattr(tok, "padding_side", "right") tok.truncation_side = "left"; tok.padding_side = "right" try: enc = tok(full_texts, return_tensors="pt", padding=True, truncation=True, max_length=self.max_length, add_special_tokens=False) input_ids = enc["input_ids"]; attention_mask = enc["attention_mask"] labels = torch.full_like(input_ids, fill_value=-100) full_ids_nt = tok(full_texts, return_tensors=None, padding=False, truncation=False, add_special_tokens=False)["input_ids"] prefix_ids_nt = tok(prefix_texts, return_tensors=None, padding=False, truncation=False, add_special_tokens=False)["input_ids"] marker_seqs = [] if MASK_COT and OUTPUT_MARKERS: for m in OUTPUT_MARKERS: mid = tok(m, add_special_tokens=False, truncation=False)["input_ids"] if not mid: continue mid_nl = tok(m+"\n", add_special_tokens=False, truncation=False)["input_ids"] marker_seqs.append((mid, mid_nl)) for i in range(input_ids.size(0)): trunc_left = max(0, len(full_ids_nt[i]) - self.max_length) boundary = len(prefix_ids_nt[i]) - trunc_left full_len_tr = int(attention_mask[i].sum().item()) if boundary <= 0 or boundary >= full_len_tr: continue span_start = boundary; span_end = full_len_tr; learn_start = span_start if MASK_COT and marker_seqs: visible_ids = input_ids[i, :full_len_tr].tolist() assistant_ids = visible_ids[span_start:span_end] best_out = None for mid, mid_nl in marker_seqs: p = self._find_subseq(assistant_ids, mid_nl) if p != -1: out_pos = span_start + p; after_pos = out_pos + len(mid_nl) else: p = self._find_subseq(assistant_ids, mid) if p == -1: continue out_pos = span_start + p; after_pos = out_pos + len(mid) if best_out is None or out_pos < best_out[0]: best_out = (out_pos, after_pos) if best_out is not None: out_pos, after_pos = best_out learn_start = after_pos if OUTPUT_LEARN_MODE != "from_marker" else out_pos learn_start = max(span_start, min(learn_start, span_end)) if learn_start < span_end: labels[i, learn_start:span_end] = input_ids[i, learn_start:span_end] labels[attention_mask == 0] = -100 return {"input_ids": input_ids, "attention_mask": attention_mask, "labels": labels} finally: tok.truncation_side = old_trunc; tok.padding_side = old_pad @torch.no_grad() def filter_has_supervision(ds, collator): keep = [] for i in range(len(ds)): out = collator([ds[i]]) if (out["labels"][0] != -100).sum().item() > 0: keep.append(i) return ds.select(keep) def count_all_masked(ds, collator, n=200, seed=3407): rng = random.Random(seed); n = min(n, len(ds)) idxs = [rng.randrange(0, len(ds)) for _ in range(n)] all_masked = 0 for i in idxs: out = collator([ds[i]]) if (out["labels"][0] != -100).sum().item() == 0: all_masked += 1 print(f"[CHECK] all-masked in {n}: {all_masked} ({all_masked/max(1,n):.1%})") def apply_upsampling(train_ds): if not UPSAMPLE_ENABLE or not UPSAMPLE_RULES_JSON: return train_ds try: rules = json.loads(UPSAMPLE_RULES_JSON) if not isinstance(rules, dict) or not rules: return train_ds except: return train_ds packs = train_ds["subcategory"] if "subcategory" in train_ds.column_names else [None]*len(train_ds) pack_field = train_ds["pack"] if "pack" in train_ds.column_names else [None]*len(train_ds) w = [] for sub, pk in zip(packs, pack_field): wt = 1.0; ss = str(sub or ""); sp = str(pk or "") for pat, mult in rules.items(): try: m = float(mult) except: m = 1.0 if pat.startswith("pack:"): if sp == pat.split(":",1)[1]: wt *= max(0.0, m) else: if pat in ss: wt *= max(0.0, m) w.append(wt) w = np.asarray(w, dtype=np.float64) if (w <= 0).all() or w.sum() == 0: return train_ds p = w / w.sum(); n = len(train_ds) idx = np.random.choice(np.arange(n), size=n, replace=True, p=p) return train_ds.select(idx.tolist()) class LabelStatsCallback(TrainerCallback): def __init__(self, dataset, collator, name="train", every_n_steps=100): self.dataset, self.collator, self.name, self.every_n_steps = dataset, collator, name, every_n_steps @torch.no_grad() def on_step_end(self, args, state, control, **kwargs): if (state.global_step % self.every_n_steps) == 0: batch = [self.dataset[random.randint(0, len(self.dataset)-1)] for _ in range(8)] out = self.collator(batch) valid = (out["labels"] != -100).sum().item() total = (out["attention_mask"] == 1).sum().item() print(f"\n[LabelStats:{self.name}] step={state.global_step} valid_ratio={valid/max(1,total):.4f}") def main(): os.makedirs(OUT_LORA_DIR, exist_ok=True) print(f"[INFO] Loading dataset: {DATASET_ID}") ds_all = load_dataset(DATASET_ID, split="train") ensure_openai_messages(ds_all) ds_all = ds_all.filter(lambda ex: has_any_nonempty_assistant_turn(ex["messages"])) ds_all = ds_all.filter(ends_with_nonempty_assistant) train_ds, val_ds = shuffle_split(ds_all, VAL_RATIO, SEED) train_ds = apply_upsampling(train_ds) print("[INFO] Loading base model:", BASE_MODEL_ID) model, tokenizer = FastLanguageModel.from_pretrained( model_name=BASE_MODEL_ID, max_seq_length=MAX_SEQ_LEN, dtype=None, load_in_4bit=True) build_cache = make_text_cache_builder(tokenizer) train_ds = train_ds.map(build_cache, batched=True, num_proc=1, desc="Caching train") val_ds = val_ds.map(build_cache, batched=True, num_proc=1, desc="Caching val") model = FastLanguageModel.get_peft_model( model, r=LORA_R, target_modules=LORA_TARGET_MODULES, lora_alpha=LORA_ALPHA, lora_dropout=LORA_DROPOUT, use_gradient_checkpointing="unsloth", random_state=SEED) args = TrainingArguments( output_dir=OUT_LORA_DIR, num_train_epochs=NUM_TRAIN_EPOCHS, per_device_train_batch_size=PER_DEVICE_TRAIN_BATCH_SIZE, per_device_eval_batch_size=PER_DEVICE_EVAL_BATCH_SIZE, gradient_accumulation_steps=GRAD_ACCUM, learning_rate=LR, warmup_ratio=WARMUP_RATIO, lr_scheduler_type="cosine", weight_decay=WEIGHT_DECAY, logging_steps=LOGGING_STEPS, eval_strategy="steps", eval_steps=EVAL_STEPS, save_strategy="steps", save_steps=SAVE_STEPS, save_total_limit=SAVE_TOTAL_LIMIT, max_steps=MAX_STEPS, bf16=False, fp16=True, push_to_hub=False, report_to="none", group_by_length=False, remove_unused_columns=False) collator = AssistantOnlyCollatorCached(tokenizer=tokenizer, max_length=MAX_SEQ_LEN) print("[INFO] Checking all-masked before filtering...") count_all_masked(val_ds, collator, n=len(val_ds), seed=SEED) print("[INFO] Filtering train/val...") train_ds = filter_has_supervision(train_ds, collator) val_ds = filter_has_supervision(val_ds, collator) print("[INFO] New sizes: train =", len(train_ds), "val =", len(val_ds)) count_all_masked(val_ds, collator, n=len(val_ds), seed=SEED) trainer = Trainer( model=model, args=args, train_dataset=train_ds, eval_dataset=val_ds, data_collator=collator, tokenizer=tokenizer) trainer.add_callback(LabelStatsCallback(train_ds, collator, name="train", every_n_steps=LOGGING_STEPS)) print("[INFO] Starting training...") trainer.train() print("[INFO] Saving adapter & tokenizer...") model.save_pretrained(OUT_LORA_DIR) tokenizer.save_pretrained(OUT_LORA_DIR) print(f"[INFO] Done. Saved to {OUT_LORA_DIR}") if __name__ == "__main__": main()