| 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() |
|
|