Text Generation
PEFT
Safetensors
Arabic
arabic
relation-extraction
qlora
bitsandbytes
multiple-choice
conversational
Instructions to use U4RASD/DRU-RE-Yehia with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use U4RASD/DRU-RE-Yehia with PEFT:
from peft import PeftModel from transformers import AutoModelForCausalLM base_model = AutoModelForCausalLM.from_pretrained("Navid-AI/Yehia-7B-preview") model = PeftModel.from_pretrained(base_model, "U4RASD/DRU-RE-Yehia") - Notebooks
- Google Colab
- Kaggle
| #!/usr/bin/env python3 | |
| """Train and evaluate DRU-RE-Yehia with one-token multiple-choice QLoRA. | |
| The model sees Arabic relation options labeled by one-token Arabic symbols and | |
| is supervised only on the selected symbol. Evaluation performs constrained | |
| greedy choice among symbols valid for each row, so formatting cannot become an | |
| invalid output. The transformed dataset is read locally. Yehia is also loaded | |
| locally; when its snapshot is absent, this entrypoint can download it once with | |
| HF_TOKENONE and then continue from the local snapshot. | |
| """ | |
| from __future__ import annotations | |
| import inspect | |
| import json | |
| import logging | |
| import math | |
| import os | |
| import random | |
| import shutil | |
| from collections import Counter, defaultdict | |
| from datetime import datetime, timezone | |
| from pathlib import Path | |
| from typing import Any, Dict, Iterable, List, Mapping, Optional, Sequence, Tuple | |
| os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") | |
| os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = "0" | |
| import bitsandbytes as bnb | |
| import numpy as np | |
| import torch | |
| from dotenv import load_dotenv | |
| from huggingface_hub import snapshot_download | |
| from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training | |
| from sklearn.metrics import accuracy_score, f1_score | |
| from torch.nn.utils.rnn import pad_sequence | |
| from torch.utils.data import Dataset, WeightedRandomSampler | |
| from tqdm.auto import tqdm | |
| from transformers import ( | |
| AutoModelForCausalLM, | |
| AutoTokenizer, | |
| BitsAndBytesConfig, | |
| Trainer, | |
| TrainerCallback, | |
| TrainingArguments, | |
| set_seed, | |
| ) | |
| from transformers.pytorch_utils import Conv1D | |
| from transformers.trainer_utils import get_last_checkpoint | |
| from re_sft_common import ( | |
| NO_RELATION_FULL, | |
| OPTION_CODES, | |
| PROMPT_VERSION, | |
| env_bool, | |
| env_float, | |
| env_int, | |
| env_str, | |
| load_jsonl, | |
| validate_local_hf_revision, | |
| ) | |
| load_dotenv() | |
| def utc_now() -> str: | |
| return datetime.now(timezone.utc).replace(microsecond=0).isoformat() | |
| def dump_json(value: Any, path: Path) -> None: | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| with path.open("w", encoding="utf-8") as handle: | |
| json.dump(value, handle, ensure_ascii=False, indent=2) | |
| handle.write("\n") | |
| def dump_jsonl(rows: Iterable[Mapping[str, Any]], path: Path) -> None: | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| with path.open("w", encoding="utf-8") as handle: | |
| for row in rows: | |
| handle.write(json.dumps(dict(row), ensure_ascii=False) + "\n") | |
| DATASET_DIR = Path(env_str("CHOICE_DATASET_DIR", "data/Yehia-RE-SFT")) | |
| BASE_MODEL_ID = env_str("YEHIA_BASE_MODEL_ID", "Navid-AI/Yehia-7B-preview") | |
| BASE_MODEL_REVISION = env_str( | |
| "YEHIA_BASE_MODEL_REVISION", "b9dda4715eafee7e8090d2c83cfe078d75f4ebb8" | |
| ) | |
| MODEL_DIR = Path(env_str("LOCAL_YEHIA_MODEL_DIR", "models/Yehia-7B-preview")) | |
| ALLOW_MODEL_DOWNLOAD = env_bool("ALLOW_MODEL_DOWNLOAD", True) | |
| OUTPUT_DIR = Path(env_str("CHOICE_OUTPUT_DIR", "runs/DRU-RE-Yehia")) | |
| CHECKPOINT_DIR = OUTPUT_DIR / "checkpoints" | |
| ARTIFACT_DIR = OUTPUT_DIR / "artifacts" | |
| METRICS_DIR = OUTPUT_DIR / "metrics" | |
| PREDICTIONS_DIR = OUTPUT_DIR / "predictions" | |
| LOG_DIR = Path(env_str("LOG_DIR", "logs")) | |
| LOG_FILE = LOG_DIR / env_str("CHOICE_TRAINING_LOG_FILE", "train.log") | |
| SEED = env_int("SEED", 42) | |
| MAX_LENGTH = env_int("MAX_LENGTH", 1024) | |
| NUM_EPOCHS = env_float("NUM_EPOCHS", 3.0) | |
| MAX_STEPS = env_int("CHOICE_MAX_STEPS", -1) | |
| TRAIN_BATCH = env_int("TRAIN_BATCH_SIZE", 4) | |
| EVAL_BATCH = env_int("CHOICE_EVAL_BATCH_SIZE", 16) | |
| GRAD_ACCUM = env_int("GRAD_ACCUM_STEPS", 4) | |
| LEARNING_RATE = env_float("LEARNING_RATE", 5e-5) | |
| WARMUP_RATIO = env_float("WARMUP_RATIO", 0.05) | |
| WEIGHT_DECAY = env_float("WEIGHT_DECAY", 0.0) | |
| MAX_GRAD_NORM = env_float("MAX_GRAD_NORM", 1.0) | |
| LOGGING_STEPS = env_int("LOGGING_STEPS", 20) | |
| EVAL_STEPS = env_int("SAVE_EVAL_STEPS", 250) | |
| SAVE_TOTAL_LIMIT = env_int("SAVE_TOTAL_LIMIT", 2) | |
| CLASS_SAMPLING_ALPHA = env_float("CLASS_SAMPLING_ALPHA", 0.25) | |
| CALIBRATE_NO_RELATION = env_bool("CALIBRATE_NO_RELATION", True) | |
| LORA_R = env_int("LORA_R", 16) | |
| LORA_ALPHA = env_int("LORA_ALPHA", 32) | |
| LORA_DROPOUT = env_float("LORA_DROPOUT", 0.05) | |
| USE_RSLORA = env_bool("USE_RSLORA", True) | |
| RAW_TARGETS = env_str("LORA_TARGET_MODULES", "all-linear").strip() | |
| TARGET_MODULES: str | List[str] = ( | |
| "all-linear" | |
| if RAW_TARGETS == "all-linear" | |
| else [item.strip() for item in RAW_TARGETS.split(",") if item.strip()] | |
| ) | |
| ATTN_IMPL = env_str("ATTENTION_IMPLEMENTATION", "sdpa") | |
| RESUME = env_bool("CHOICE_RESUME", True) | |
| for path in (OUTPUT_DIR, CHECKPOINT_DIR, ARTIFACT_DIR, METRICS_DIR, PREDICTIONS_DIR, LOG_DIR): | |
| path.mkdir(parents=True, exist_ok=True) | |
| logging.basicConfig( | |
| level=logging.INFO, | |
| format="%(asctime)s | %(levelname)s | %(message)s", | |
| handlers=[logging.FileHandler(LOG_FILE, encoding="utf-8"), logging.StreamHandler()], | |
| ) | |
| logger = logging.getLogger("yehia_choice_qlora") | |
| def stage(name: str) -> None: | |
| message = f"\n{'#' * 96}\n# {name}\n{'#' * 96}" | |
| print(message, flush=True) | |
| logger.info(message) | |
| set_seed(SEED) | |
| random.seed(SEED) | |
| np.random.seed(SEED) | |
| stage("CONFIGURATION") | |
| config = { | |
| "created_at_utc": utc_now(), | |
| "dataset_dir": str(DATASET_DIR), | |
| "base_model_id": BASE_MODEL_ID, | |
| "base_model_revision": BASE_MODEL_REVISION, | |
| "model_dir": str(MODEL_DIR), | |
| "allow_model_download": ALLOW_MODEL_DOWNLOAD, | |
| "output_dir": str(OUTPUT_DIR), | |
| "prompt_version": PROMPT_VERSION, | |
| "objective": "one_token_full_vocab_multiple_choice", | |
| "loss_mode": "full_vocab", | |
| "decision_tokens": list(OPTION_CODES), | |
| "max_length": MAX_LENGTH, | |
| "num_epochs": NUM_EPOCHS, | |
| "max_steps": MAX_STEPS, | |
| "train_batch": TRAIN_BATCH, | |
| "eval_batch": EVAL_BATCH, | |
| "gradient_accumulation": GRAD_ACCUM, | |
| "effective_batch": TRAIN_BATCH * GRAD_ACCUM, | |
| "learning_rate": LEARNING_RATE, | |
| "warmup_ratio": WARMUP_RATIO, | |
| "weight_decay": WEIGHT_DECAY, | |
| "max_grad_norm": MAX_GRAD_NORM, | |
| "eval_steps": EVAL_STEPS, | |
| "class_sampling_alpha": CLASS_SAMPLING_ALPHA, | |
| "calibrate_no_relation": CALIBRATE_NO_RELATION, | |
| "lora_r": LORA_R, | |
| "lora_alpha": LORA_ALPHA, | |
| "lora_dropout": LORA_DROPOUT, | |
| "use_rslora": USE_RSLORA, | |
| "target_modules": TARGET_MODULES, | |
| "seed": SEED, | |
| } | |
| dump_json(config, OUTPUT_DIR / "run_config.json") | |
| print(json.dumps(config, ensure_ascii=False, indent=2), flush=True) | |
| stage("LOAD AND VALIDATE LOCAL DATASET") | |
| required_files = { | |
| "train": DATASET_DIR / "train.jsonl", | |
| "validation": DATASET_DIR / "validation.jsonl", | |
| "official": DATASET_DIR / "official.jsonl", | |
| } | |
| for path in required_files.values(): | |
| if not path.exists(): | |
| raise FileNotFoundError(path) | |
| rows = {split: load_jsonl(path) for split, path in required_files.items()} | |
| expected_counts = {"train": 15686, "validation": 1687, "official": 2074} | |
| for split, expected in expected_counts.items(): | |
| if len(rows[split]) != expected: | |
| raise ValueError(f"{split} count {len(rows[split])} != {expected}") | |
| for split in ("train", "validation"): | |
| for index, row in enumerate(rows[split]): | |
| options = row.get("allowed_options_ar") or [] | |
| codes = row.get("option_codes") or [] | |
| gold_index = row.get("gold_option_index") | |
| gold_code = row.get("gold_answer_code") | |
| if row.get("prompt_version") != PROMPT_VERSION: | |
| raise ValueError(f"Prompt mismatch at {split}:{index}") | |
| if not options or options[-1] != "لا توجد علاقة": | |
| raise ValueError(f"No-relation is not last at {split}:{index}") | |
| if codes != list(OPTION_CODES[: len(options)]): | |
| raise ValueError(f"Option-code mismatch at {split}:{index}") | |
| if not isinstance(gold_index, int) or not 0 <= gold_index < len(options): | |
| raise ValueError(f"Gold option index mismatch at {split}:{index}") | |
| if gold_code != codes[gold_index]: | |
| raise ValueError(f"Gold code mismatch at {split}:{index}") | |
| if row["messages"][-1] != {"role": "assistant", "content": gold_code}: | |
| raise ValueError(f"Assistant target mismatch at {split}:{index}") | |
| for index, row in enumerate(rows["official"]): | |
| if row.get("gold_answer_code") is not None or len(row.get("messages", [])) != 2: | |
| raise ValueError(f"Official label leakage at row {index}") | |
| all_labels = sorted( | |
| {str(row["gold_relation_full"]) for row in rows["train"] + rows["validation"]} | |
| ) | |
| if len(all_labels) != 41 or NO_RELATION_FULL not in all_labels: | |
| raise ValueError(f"Expected 41 labels; found {len(all_labels)}") | |
| class_counts = Counter(str(row["gold_relation_full"]) for row in rows["train"]) | |
| dump_json(dict(sorted(class_counts.items())), ARTIFACT_DIR / "train_class_counts.json") | |
| print({split: len(value) for split, value in rows.items()}, flush=True) | |
| def require_yehia_token() -> str: | |
| """Read the gated-model token without ever logging its value.""" | |
| token = os.environ.get("HF_TOKENONE", "").strip() | |
| if not token or token.startswith("hf_your_"): | |
| raise RuntimeError( | |
| "Yehia is not available locally and HF_TOKENONE is missing. " | |
| "Add the gated-model token to .env or pre-populate LOCAL_YEHIA_MODEL_DIR." | |
| ) | |
| return token | |
| def ensure_local_model_snapshot() -> None: | |
| """Use an existing local snapshot or download Yehia into MODEL_DIR once.""" | |
| if (MODEL_DIR / "config.json").is_file(): | |
| revision = validate_local_hf_revision(MODEL_DIR, BASE_MODEL_REVISION) | |
| print(f"Using existing local Yehia snapshot: {MODEL_DIR}", flush=True) | |
| print(f"Verified local Yehia revision: {revision}", flush=True) | |
| return | |
| if not ALLOW_MODEL_DOWNLOAD: | |
| raise FileNotFoundError( | |
| f"Local Yehia snapshot is missing at {MODEL_DIR} and " | |
| "ALLOW_MODEL_DOWNLOAD=false" | |
| ) | |
| offline = os.environ.get("HF_HUB_OFFLINE", "").strip().lower() | |
| if offline in {"1", "true", "yes", "on"}: | |
| raise RuntimeError( | |
| "Local Yehia snapshot is missing but HF_HUB_OFFLINE is enabled. " | |
| "Disable offline mode for the first run or download the model separately." | |
| ) | |
| MODEL_DIR.mkdir(parents=True, exist_ok=True) | |
| print(f"Downloading gated base model {BASE_MODEL_ID} into {MODEL_DIR}", flush=True) | |
| snapshot_download( | |
| repo_id=BASE_MODEL_ID, | |
| revision=BASE_MODEL_REVISION, | |
| local_dir=str(MODEL_DIR), | |
| token=require_yehia_token(), | |
| ) | |
| if not (MODEL_DIR / "config.json").is_file(): | |
| raise RuntimeError(f"Downloaded Yehia snapshot is incomplete: {MODEL_DIR}") | |
| (MODEL_DIR / ".dru_hf_revision").write_text(BASE_MODEL_REVISION + "\n", encoding="utf-8") | |
| validate_local_hf_revision(MODEL_DIR, BASE_MODEL_REVISION) | |
| stage("RESOLVE LOCAL YEHIA SNAPSHOT AND BUILD DECISION-TOKEN LABELS") | |
| ensure_local_model_snapshot() | |
| tokenizer = AutoTokenizer.from_pretrained( | |
| str(MODEL_DIR), local_files_only=True, use_fast=True | |
| ) | |
| if not tokenizer.chat_template: | |
| raise RuntimeError("Local Yehia tokenizer has no chat template") | |
| if tokenizer.pad_token_id is None: | |
| tokenizer.pad_token = tokenizer.eos_token | |
| tokenizer.padding_side = "right" | |
| code_token_ids: List[int] = [] | |
| for code in OPTION_CODES: | |
| ids = tokenizer.encode(" " + code, add_special_tokens=False) | |
| if len(ids) != 1: | |
| raise RuntimeError(f"Decision code {code!r} is not one token: {ids}") | |
| code_token_ids.append(int(ids[0])) | |
| dump_json( | |
| {"codes": list(OPTION_CODES), "token_ids": code_token_ids}, | |
| ARTIFACT_DIR / "decision_token_inventory.json", | |
| ) | |
| def encode_row(row: Mapping[str, Any]) -> Dict[str, Any]: | |
| prompt_ids = list( | |
| tokenizer.apply_chat_template( | |
| row["prompt_messages"], tokenize=True, add_generation_prompt=True | |
| ) | |
| ) | |
| full_ids = list( | |
| tokenizer.apply_chat_template( | |
| row["messages"], tokenize=True, add_generation_prompt=False | |
| ) | |
| ) | |
| common = 0 | |
| for first, second in zip(prompt_ids, full_ids): | |
| if first != second: | |
| break | |
| common += 1 | |
| if common != len(prompt_ids): | |
| raise ValueError( | |
| f"Prompt/full boundary mismatch for {row.get('id')}: " | |
| f"prompt={len(prompt_ids)}, common={common}" | |
| ) | |
| if len(full_ids) > MAX_LENGTH: | |
| raise ValueError(f"Sequence exceeds {MAX_LENGTH} tokens for {row.get('id')}") | |
| expected_token = code_token_ids[int(row["gold_option_index"])] | |
| if common >= len(full_ids) or full_ids[common] != expected_token: | |
| raise ValueError( | |
| f"Decision token mismatch for {row.get('id')}: " | |
| f"expected={expected_token}, observed={full_ids[common] if common < len(full_ids) else None}" | |
| ) | |
| # Pure classification objective: only the one decision token contributes. | |
| labels = [-100] * len(full_ids) | |
| labels[common] = expected_token | |
| return { | |
| "input_ids": full_ids, | |
| "attention_mask": [1] * len(full_ids), | |
| "labels": labels, | |
| "prompt_token_count": common, | |
| "decision_token_index": common, | |
| "decision_token_id": expected_token, | |
| } | |
| encoded = { | |
| split: [encode_row(row) for row in tqdm(rows[split], desc=f"tokenize:{split}")] | |
| for split in ("train", "validation") | |
| } | |
| for split, values in encoded.items(): | |
| lengths = [len(item["input_ids"]) for item in values] | |
| print( | |
| f"{split}: min={min(lengths)} mean={np.mean(lengths):.1f} " | |
| f"p95={np.percentile(lengths, 95):.1f} max={max(lengths)}", | |
| flush=True, | |
| ) | |
| audit_indices = [ | |
| 0, | |
| next(i for i, row in enumerate(rows["train"]) if row["gold_relation_full"] == NO_RELATION_FULL), | |
| max(range(len(encoded["train"])), key=lambda i: len(encoded["train"][i]["input_ids"])), | |
| ] | |
| masking_audit: List[Dict[str, Any]] = [] | |
| for index in dict.fromkeys(audit_indices): | |
| row = rows["train"][index] | |
| item = encoded["train"][index] | |
| supervised_positions = [i for i, label in enumerate(item["labels"]) if label != -100] | |
| if supervised_positions != [item["decision_token_index"]]: | |
| raise RuntimeError(f"Non-decision tokens are supervised for train row {index}") | |
| masking_audit.append( | |
| { | |
| "index": index, | |
| "id": row["id"], | |
| "gold_relation_full": row["gold_relation_full"], | |
| "gold_answer_code": row["gold_answer_code"], | |
| "decoded_conversation": tokenizer.decode( | |
| item["input_ids"], skip_special_tokens=False | |
| ), | |
| "prompt_token_count": item["prompt_token_count"], | |
| "masked_prompt_token_count": item["prompt_token_count"], | |
| "supervised_token_count": 1, | |
| "supervised_token_id": item["decision_token_id"], | |
| "decoded_supervised_token": tokenizer.decode([item["decision_token_id"]]), | |
| "all_nondecision_tokens_masked": True, | |
| } | |
| ) | |
| dump_json( | |
| {"status": "passed", "objective": "decision_token_only", "rows": masking_audit}, | |
| ARTIFACT_DIR / "decision_token_masking_audit.json", | |
| ) | |
| class EncodedDataset(Dataset): | |
| def __init__(self, items: Sequence[Mapping[str, Any]]): | |
| self.items = list(items) | |
| def __len__(self) -> int: | |
| return len(self.items) | |
| def __getitem__(self, index: int) -> Dict[str, torch.Tensor]: | |
| item = self.items[index] | |
| return { | |
| key: torch.tensor(item[key], dtype=torch.long) | |
| for key in ("input_ids", "attention_mask", "labels") | |
| } | |
| class CompletionCollator: | |
| def __call__( | |
| self, features: Sequence[Mapping[str, torch.Tensor]] | |
| ) -> Dict[str, torch.Tensor]: | |
| return { | |
| "input_ids": pad_sequence( | |
| [item["input_ids"] for item in features], | |
| batch_first=True, | |
| padding_value=tokenizer.pad_token_id, | |
| ), | |
| "attention_mask": pad_sequence( | |
| [item["attention_mask"] for item in features], | |
| batch_first=True, | |
| padding_value=0, | |
| ), | |
| "labels": pad_sequence( | |
| [item["labels"] for item in features], | |
| batch_first=True, | |
| padding_value=-100, | |
| ), | |
| } | |
| train_dataset = EncodedDataset(encoded["train"]) | |
| validation_dataset = EncodedDataset(encoded["validation"]) | |
| collator = CompletionCollator() | |
| def normalize_module_name(name: str) -> str: | |
| prefixes = ("base_model.model.", "model.", "module.") | |
| changed = True | |
| while changed: | |
| changed = False | |
| for prefix in prefixes: | |
| if name.startswith(prefix): | |
| name = name[len(prefix) :] | |
| changed = True | |
| return name | |
| def supported_linear_classes() -> Tuple[type, ...]: | |
| classes: List[type] = [torch.nn.Linear, Conv1D] | |
| for attr in ("Linear4bit", "Linear8bitLt"): | |
| cls = getattr(getattr(bnb, "nn", object()), attr, None) | |
| if isinstance(cls, type): | |
| classes.append(cls) | |
| return tuple(classes) | |
| def module_class_name(module: torch.nn.Module) -> str: | |
| return f"{module.__class__.__module__}.{module.__class__.__name__}" | |
| def enumerate_eligible_linear_modules(active_model: torch.nn.Module) -> Dict[str, Any]: | |
| """Inventory every PEFT-supported projection before attaching adapters.""" | |
| supported = supported_linear_classes() | |
| output_embedding = active_model.get_output_embeddings() | |
| input_embedding = active_model.get_input_embeddings() | |
| output_ids = {id(output_embedding)} if output_embedding is not None else set() | |
| input_ids = {id(input_embedding)} if input_embedding is not None else set() | |
| records: List[Dict[str, Any]] = [] | |
| output_head_names: List[str] = [] | |
| embedding_names: List[str] = [] | |
| for name, module in active_model.named_modules(): | |
| if not name: | |
| continue | |
| normalized = normalize_module_name(name) | |
| is_output_head = id(module) in output_ids | |
| is_embedding = id(module) in input_ids or isinstance(module, torch.nn.Embedding) | |
| if is_output_head: | |
| output_head_names.append(normalized) | |
| if is_embedding: | |
| embedding_names.append(normalized) | |
| if isinstance(module, supported): | |
| records.append( | |
| { | |
| "name": normalized, | |
| "original_name": name, | |
| "class": module_class_name(module), | |
| "suffix": normalized.rsplit(".", 1)[-1], | |
| "is_output_head": bool(is_output_head), | |
| "is_embedding": bool(is_embedding), | |
| } | |
| ) | |
| expected_wrapped = [ | |
| item["name"] | |
| for item in records | |
| if not item["is_output_head"] and not item["is_embedding"] | |
| ] | |
| inventory = { | |
| "target_modules": TARGET_MODULES, | |
| "supported_classes": [f"{cls.__module__}.{cls.__name__}" for cls in supported], | |
| "total_eligible_linear_or_conv1d_modules": len(records), | |
| "expected_lora_wrapped_modules": len(expected_wrapped), | |
| "output_head_names": sorted(output_head_names), | |
| "embedding_names": sorted(embedding_names), | |
| "modules": sorted(records, key=lambda item: item["name"]), | |
| "expected_wrapped_module_names": sorted(expected_wrapped), | |
| } | |
| dump_json(inventory, ARTIFACT_DIR / "all_linear_eligible_modules.json") | |
| return inventory | |
| def trainable_parameter_summary(active_model: torch.nn.Module) -> Dict[str, Any]: | |
| if hasattr(active_model, "get_nb_trainable_parameters"): | |
| trainable, total = active_model.get_nb_trainable_parameters() | |
| else: | |
| total = sum(parameter.numel() for parameter in active_model.parameters()) | |
| trainable = sum( | |
| parameter.numel() | |
| for parameter in active_model.parameters() | |
| if parameter.requires_grad | |
| ) | |
| trainable_names = [ | |
| name for name, parameter in active_model.named_parameters() if parameter.requires_grad | |
| ] | |
| unintended = [name for name in trainable_names if "lora_" not in name] | |
| if unintended: | |
| raise RuntimeError(f"Unexpected non-LoRA trainable parameters: {unintended[:20]}") | |
| return { | |
| "total_parameters": int(total), | |
| "trainable_parameters": int(trainable), | |
| "trainable_percentage": float(100.0 * trainable / total if total else 0.0), | |
| "trainable_tensor_count": len(trainable_names), | |
| "trainable_parameter_names": trainable_names, | |
| } | |
| def validate_lora_wrapping( | |
| active_model: torch.nn.Module, eligible_inventory: Mapping[str, Any] | |
| ) -> Dict[str, Any]: | |
| """Prove true all-linear coverage and that the frozen head/base stayed frozen.""" | |
| wrapped: List[Dict[str, Any]] = [] | |
| for name, module in active_model.named_modules(): | |
| if hasattr(module, "lora_A") and hasattr(module, "lora_B"): | |
| normalized = normalize_module_name(name) | |
| wrapped.append( | |
| { | |
| "name": normalized, | |
| "original_name": name, | |
| "class": module_class_name(module), | |
| "suffix": normalized.rsplit(".", 1)[-1], | |
| } | |
| ) | |
| wrapped_names = {item["name"] for item in wrapped} | |
| expected_names = set(eligible_inventory["expected_wrapped_module_names"]) | |
| output_head_names = set(eligible_inventory["output_head_names"]) | |
| if TARGET_MODULES == "all-linear": | |
| missing = sorted(expected_names - wrapped_names) | |
| unexpected = sorted((wrapped_names - expected_names) - output_head_names) | |
| if missing: | |
| raise RuntimeError(f"Incomplete all-linear LoRA coverage: {missing[:50]}") | |
| if unexpected: | |
| raise RuntimeError(f"Unexpected all-linear LoRA modules: {unexpected[:50]}") | |
| if wrapped_names & output_head_names: | |
| raise RuntimeError( | |
| "Output head was unintentionally adapted: " | |
| f"{sorted(wrapped_names & output_head_names)}" | |
| ) | |
| required = {"q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"} | |
| wrapped_suffixes = {item["suffix"] for item in wrapped} | |
| missing_required = sorted(required - wrapped_suffixes) | |
| if missing_required: | |
| raise RuntimeError(f"Missing attention/MLP adapters: {missing_required}") | |
| by_suffix = Counter(item["suffix"] for item in wrapped) | |
| by_layer: Counter[str] = Counter() | |
| for item in wrapped: | |
| parts = item["name"].split(".") | |
| layer = next( | |
| ( | |
| f"{parts[index]}.{parts[index + 1]}" | |
| for index in range(len(parts) - 1) | |
| if parts[index] == "layers" | |
| ), | |
| "non_layer", | |
| ) | |
| by_layer[layer] += 1 | |
| audit = { | |
| "target_modules": TARGET_MODULES, | |
| "eligible_linear_or_conv1d_count": eligible_inventory[ | |
| "total_eligible_linear_or_conv1d_modules" | |
| ], | |
| "expected_lora_wrapped_count": len(expected_names), | |
| "lora_wrapped_count": len(wrapped), | |
| "wrapped_by_suffix": dict(sorted(by_suffix.items())), | |
| "wrapped_by_transformer_layer": dict(sorted(by_layer.items())), | |
| "output_head_names": sorted(output_head_names), | |
| "output_head_adapted": False, | |
| "attention_projection_adapters_present": all( | |
| suffix in wrapped_suffixes for suffix in ("q_proj", "k_proj", "v_proj", "o_proj") | |
| ), | |
| "mlp_projection_adapters_present": all( | |
| suffix in wrapped_suffixes for suffix in ("gate_proj", "up_proj", "down_proj") | |
| ), | |
| "base_quantized_weights_trainable": False, | |
| "only_lora_parameters_trainable": True, | |
| "wrapped_modules": sorted(wrapped, key=lambda item: item["name"]), | |
| **trainable_parameter_summary(active_model), | |
| } | |
| dump_json(audit, ARTIFACT_DIR / "lora_wrapped_modules.json") | |
| return audit | |
| def positive_scores(gold: Sequence[str], pred: Sequence[str]) -> Dict[str, Any]: | |
| tp = fp = fn = 0 | |
| for truth, guess in zip(gold, pred): | |
| truth_positive = truth != NO_RELATION_FULL | |
| guess_positive = guess != NO_RELATION_FULL | |
| if truth_positive and guess_positive and truth == guess: | |
| tp += 1 | |
| else: | |
| fp += int(guess_positive) | |
| fn += int(truth_positive) | |
| precision = tp / (tp + fp) if tp + fp else 0.0 | |
| recall = tp / (tp + fn) if tp + fn else 0.0 | |
| f1 = 2 * precision * recall / (precision + recall) if precision + recall else 0.0 | |
| return { | |
| "positive_micro_precision": precision, | |
| "positive_micro_recall": recall, | |
| "positive_micro_f1": f1, | |
| "positive_tp": tp, | |
| "positive_fp": fp, | |
| "positive_fn": fn, | |
| } | |
| def relation_metrics(gold: Sequence[str], pred: Sequence[str]) -> Dict[str, Any]: | |
| return { | |
| "rows": len(gold), | |
| "accuracy": float(accuracy_score(gold, pred)), | |
| "macro_f1_41": float( | |
| f1_score(gold, pred, labels=all_labels, average="macro", zero_division=0) | |
| ), | |
| "weighted_f1_41": float( | |
| f1_score(gold, pred, labels=all_labels, average="weighted", zero_division=0) | |
| ), | |
| **positive_scores(gold, pred), | |
| "invalid_outputs": 0, | |
| "invalid_output_rate": 0.0, | |
| } | |
| def choose_from_scores( | |
| records: Sequence[Mapping[str, Any]], score_rows: Sequence[Sequence[float]], bias: float | |
| ) -> Tuple[List[str], List[int]]: | |
| predictions: List[str] = [] | |
| positions: List[int] = [] | |
| for row, raw_scores in zip(records, score_rows): | |
| scores = list(raw_scores) | |
| scores[-1] += bias | |
| position = int(np.argmax(scores)) | |
| positions.append(position) | |
| predictions.append(str(row["allowed_relation_full_labels"][position])) | |
| return predictions, positions | |
| def calibrate_bias( | |
| records: Sequence[Mapping[str, Any]], score_rows: Sequence[Sequence[float]] | |
| ) -> Tuple[float, Dict[str, Any]]: | |
| gold = [str(row["gold_relation_full"]) for row in records] | |
| if not CALIBRATE_NO_RELATION: | |
| pred, _ = choose_from_scores(records, score_rows, 0.0) | |
| return 0.0, relation_metrics(gold, pred) | |
| best_bias = 0.0 | |
| best_metrics: Optional[Dict[str, Any]] = None | |
| for bias in np.arange(-12.0, 12.0001, 0.1): | |
| pred, _ = choose_from_scores(records, score_rows, float(bias)) | |
| current = relation_metrics(gold, pred) | |
| key = (current["positive_micro_f1"], current["accuracy"], -abs(float(bias))) | |
| old_key = ( | |
| (-1.0, -1.0, -math.inf) | |
| if best_metrics is None | |
| else ( | |
| best_metrics["positive_micro_f1"], | |
| best_metrics["accuracy"], | |
| -abs(best_bias), | |
| ) | |
| ) | |
| if key > old_key: | |
| best_bias, best_metrics = float(bias), current | |
| assert best_metrics is not None | |
| for bias in np.arange(best_bias - 0.1, best_bias + 0.1001, 0.01): | |
| pred, _ = choose_from_scores(records, score_rows, float(bias)) | |
| current = relation_metrics(gold, pred) | |
| key = (current["positive_micro_f1"], current["accuracy"], -abs(float(bias))) | |
| old_key = ( | |
| best_metrics["positive_micro_f1"], | |
| best_metrics["accuracy"], | |
| -abs(best_bias), | |
| ) | |
| if key > old_key: | |
| best_bias, best_metrics = float(bias), current | |
| return best_bias, best_metrics | |
| def constrained_choice_evaluation( | |
| active_model: torch.nn.Module, | |
| records: Sequence[Mapping[str, Any]], | |
| batch_size: int, | |
| step: int, | |
| prefix: str, | |
| ) -> Dict[str, Any]: | |
| old_padding = tokenizer.padding_side | |
| was_training = active_model.training | |
| tokenizer.padding_side = "left" | |
| active_model.eval() | |
| # Length bucketing preserves row order in the artifacts while avoiding the | |
| # large padding penalty caused by mixing short and long Arabic prompts. | |
| indexed_prompts: List[Tuple[int, Mapping[str, Any], List[int]]] = [] | |
| for original_index, row in enumerate(records): | |
| prompt_ids = list( | |
| tokenizer.apply_chat_template( | |
| row["prompt_messages"], tokenize=True, add_generation_prompt=True | |
| ) | |
| ) | |
| if len(prompt_ids) > MAX_LENGTH: | |
| raise ValueError("Evaluation prompt exceeds max length") | |
| indexed_prompts.append((original_index, row, prompt_ids)) | |
| indexed_prompts.sort(key=lambda item: len(item[2])) | |
| score_rows_by_index: List[Optional[List[float]]] = [None] * len(records) | |
| generated_token_ids_by_index: List[Optional[int]] = [None] * len(records) | |
| gold_nll: List[float] = [] | |
| try: | |
| for start in tqdm( | |
| range(0, len(indexed_prompts), batch_size), | |
| desc=f"choice-eval:{prefix}", | |
| leave=False, | |
| ): | |
| current = indexed_prompts[start : start + batch_size] | |
| batch = tokenizer.pad( | |
| {"input_ids": [item[2] for item in current]}, | |
| padding=True, | |
| return_tensors="pt", | |
| ) | |
| batch = {key: value.to(active_model.device) for key, value in batch.items()} | |
| with torch.inference_mode(): | |
| decoder, lm_head = decoder_and_lm_head(active_model) | |
| decoder_output = decoder( | |
| **batch, use_cache=False, return_dict=True | |
| ) | |
| decision_hidden = decoder_output.last_hidden_state[:, -1, :] | |
| full_next_logits = lm_head(decision_hidden).float() | |
| next_code_logits = full_next_logits.index_select( | |
| 1, torch.as_tensor(code_token_ids, device=decision_hidden.device) | |
| ) | |
| generated_token_ids = full_next_logits.argmax(dim=-1) | |
| for batch_index, (original_index, row, _) in enumerate(current): | |
| option_count = len(row["option_codes"]) | |
| candidates = next_code_logits[batch_index, :option_count] | |
| score_rows_by_index[original_index] = candidates.detach().cpu().tolist() | |
| generated_token_ids_by_index[original_index] = int( | |
| generated_token_ids[batch_index].detach().cpu() | |
| ) | |
| if row.get("gold_option_index") is not None: | |
| gold_position = int(row["gold_option_index"]) | |
| gold_nll.append( | |
| float( | |
| ( | |
| torch.logsumexp(candidates, dim=-1) | |
| - candidates[gold_position] | |
| ).cpu() | |
| ) | |
| ) | |
| finally: | |
| tokenizer.padding_side = old_padding | |
| if was_training: | |
| active_model.train() | |
| if any(scores is None for scores in score_rows_by_index) or any( | |
| token_id is None for token_id in generated_token_ids_by_index | |
| ): | |
| raise RuntimeError("Constrained evaluation failed to score every row") | |
| score_rows = [scores for scores in score_rows_by_index if scores is not None] | |
| generated_token_ids_complete = [ | |
| int(token_id) for token_id in generated_token_ids_by_index if token_id is not None | |
| ] | |
| gold = [str(row["gold_relation_full"]) for row in records] | |
| code_id_to_position = {token_id: index for index, token_id in enumerate(code_token_ids)} | |
| generated_pred: List[str] = [] | |
| generated_valid: List[bool] = [] | |
| for row, generated_token_id in zip(records, generated_token_ids_complete): | |
| position = code_id_to_position.get(generated_token_id) | |
| valid = position is not None and position < len(row["option_codes"]) | |
| generated_valid.append(valid) | |
| generated_pred.append( | |
| str(row["allowed_relation_full_labels"][position]) | |
| if valid and position is not None | |
| else NO_RELATION_FULL | |
| ) | |
| generated_metrics = relation_metrics(gold, generated_pred) | |
| generated_metrics["invalid_outputs"] = int(sum(not valid for valid in generated_valid)) | |
| generated_metrics["invalid_output_rate"] = float( | |
| generated_metrics["invalid_outputs"] / len(records) | |
| ) | |
| raw_pred, raw_positions = choose_from_scores(records, score_rows, 0.0) | |
| raw_metrics = relation_metrics(gold, raw_pred) | |
| bias, calibrated_metrics = calibrate_bias(records, score_rows) | |
| final_pred, final_positions = choose_from_scores(records, score_rows, bias) | |
| predictions: List[Dict[str, Any]] = [] | |
| for row, raw_scores, raw_position, final_position, generated_token_id, generated_is_valid in zip( | |
| records, | |
| score_rows, | |
| raw_positions, | |
| final_positions, | |
| generated_token_ids_complete, | |
| generated_valid, | |
| ): | |
| predictions.append( | |
| { | |
| "id": row["id"], | |
| "triple_id": row.get("triple_id"), | |
| "gold_relation_full": row.get("gold_relation_full"), | |
| "predicted_relation_full": row["allowed_relation_full_labels"][final_position], | |
| "predicted_code": row["option_codes"][final_position], | |
| "predicted_answer_ar": row["allowed_options_ar"][final_position], | |
| "raw_predicted_relation_full": row["allowed_relation_full_labels"][raw_position], | |
| "raw_predicted_code": row["option_codes"][raw_position], | |
| "candidate_logits": raw_scores, | |
| "no_relation_logit_bias": bias, | |
| "allowed_relation_full_labels": row["allowed_relation_full_labels"], | |
| "option_codes": row["option_codes"], | |
| "allowed_options_ar": row["allowed_options_ar"], | |
| "unconstrained_generated_token_id": generated_token_id, | |
| "unconstrained_generated_token": tokenizer.decode([generated_token_id]), | |
| "unconstrained_generated_valid": generated_is_valid, | |
| } | |
| ) | |
| prediction_path = PREDICTIONS_DIR / f"{prefix}_choice_step_{step}.jsonl" | |
| dump_jsonl(predictions, prediction_path) | |
| result = { | |
| **calibrated_metrics, | |
| "decision_nll": float(np.mean(gold_nll)), | |
| "raw_accuracy": raw_metrics["accuracy"], | |
| "raw_macro_f1_41": raw_metrics["macro_f1_41"], | |
| "raw_weighted_f1_41": raw_metrics["weighted_f1_41"], | |
| "raw_positive_micro_precision": raw_metrics["positive_micro_precision"], | |
| "raw_positive_micro_recall": raw_metrics["positive_micro_recall"], | |
| "raw_positive_micro_f1": raw_metrics["positive_micro_f1"], | |
| "unconstrained_accuracy": generated_metrics["accuracy"], | |
| "unconstrained_macro_f1_41": generated_metrics["macro_f1_41"], | |
| "unconstrained_weighted_f1_41": generated_metrics["weighted_f1_41"], | |
| "unconstrained_positive_micro_precision": generated_metrics[ | |
| "positive_micro_precision" | |
| ], | |
| "unconstrained_positive_micro_recall": generated_metrics[ | |
| "positive_micro_recall" | |
| ], | |
| "unconstrained_positive_micro_f1": generated_metrics["positive_micro_f1"], | |
| "unconstrained_invalid_outputs": generated_metrics["invalid_outputs"], | |
| "unconstrained_invalid_output_rate": generated_metrics["invalid_output_rate"], | |
| "no_relation_logit_bias": bias, | |
| "global_step": step, | |
| "evaluated_rows": len(records), | |
| "inference": "constrained_one_token_greedy", | |
| "prediction_file": str(prediction_path), | |
| } | |
| dump_json(result, METRICS_DIR / f"{prefix}_choice_metrics_step_{step}.json") | |
| return result | |
| stage("LOAD 4-BIT LOCAL YEHIA AND ATTACH QLORA") | |
| quantization = BitsAndBytesConfig( | |
| load_in_4bit=True, | |
| bnb_4bit_quant_type="nf4", | |
| bnb_4bit_use_double_quant=True, | |
| bnb_4bit_compute_dtype=torch.bfloat16, | |
| ) | |
| model = AutoModelForCausalLM.from_pretrained( | |
| str(MODEL_DIR), | |
| local_files_only=True, | |
| quantization_config=quantization, | |
| torch_dtype=torch.bfloat16, | |
| device_map={"": 0}, | |
| attn_implementation=ATTN_IMPL, | |
| ) | |
| model.config.use_cache = False | |
| model = prepare_model_for_kbit_training(model, use_gradient_checkpointing=True) | |
| try: | |
| model.gradient_checkpointing_enable(gradient_checkpointing_kwargs={"use_reentrant": False}) | |
| except TypeError: | |
| model.gradient_checkpointing_enable() | |
| eligible_linear_inventory = enumerate_eligible_linear_modules(model) | |
| model = get_peft_model( | |
| model, | |
| LoraConfig( | |
| r=LORA_R, | |
| lora_alpha=LORA_ALPHA, | |
| lora_dropout=LORA_DROPOUT, | |
| bias="none", | |
| task_type="CAUSAL_LM", | |
| target_modules=TARGET_MODULES, | |
| use_rslora=USE_RSLORA, | |
| ), | |
| ) | |
| model.print_trainable_parameters() | |
| lora_audit = validate_lora_wrapping(model, eligible_linear_inventory) | |
| print( | |
| json.dumps( | |
| { | |
| "eligible_linear_or_conv1d": lora_audit["eligible_linear_or_conv1d_count"], | |
| "lora_wrapped": lora_audit["lora_wrapped_count"], | |
| "trainable_parameters": lora_audit["trainable_parameters"], | |
| "trainable_percentage": lora_audit["trainable_percentage"], | |
| }, | |
| indent=2, | |
| ), | |
| flush=True, | |
| ) | |
| def decoder_and_lm_head(active_model: torch.nn.Module): | |
| """Resolve the adapted decoder and frozen output projection from PEFT.""" | |
| causal_model = ( | |
| active_model.get_base_model() | |
| if hasattr(active_model, "get_base_model") | |
| else active_model | |
| ) | |
| if not hasattr(causal_model, "model") or not hasattr(causal_model, "lm_head"): | |
| raise RuntimeError(f"Unexpected causal model structure: {type(causal_model)}") | |
| return causal_model.model, causal_model.lm_head | |
| stage("FINITE FORWARD/BACKWARD SMOKE TEST") | |
| long_indices = sorted( | |
| range(len(train_dataset)), | |
| key=lambda index: len(encoded["train"][index]["input_ids"]), | |
| reverse=True, | |
| )[:TRAIN_BATCH] | |
| torch.cuda.empty_cache() | |
| torch.cuda.reset_peak_memory_stats() | |
| smoke = collator([train_dataset[index] for index in long_indices]) | |
| if any(int((item["labels"] != -100).sum()) != 1 for item in [train_dataset[i] for i in long_indices]): | |
| raise RuntimeError("Smoke rows do not contain exactly one supervised token") | |
| smoke = {key: value.to(model.device) for key, value in smoke.items()} | |
| model.train() | |
| smoke_output = model( | |
| input_ids=smoke["input_ids"], | |
| attention_mask=smoke["attention_mask"], | |
| labels=smoke["labels"], | |
| ) | |
| smoke_loss = smoke_output.loss | |
| if not torch.isfinite(smoke_loss): | |
| raise RuntimeError(f"Non-finite smoke loss: {smoke_loss}") | |
| (smoke_loss / GRAD_ACCUM).backward() | |
| finite_nonzero_grads = 0 | |
| for name, parameter in model.named_parameters(): | |
| if parameter.grad is not None: | |
| if "lora_" not in name: | |
| raise RuntimeError(f"Frozen base parameter received a gradient: {name}") | |
| if torch.isfinite(parameter.grad).all() and float(parameter.grad.abs().sum()) > 0: | |
| finite_nonzero_grads += 1 | |
| if finite_nonzero_grads < 4: | |
| raise RuntimeError("Too few finite nonzero LoRA gradients") | |
| smoke_artifact = { | |
| "status": "passed", | |
| "loss": float(smoke_loss.detach().cpu()), | |
| "loss_mode": "full_vocab", | |
| "finite_nonzero_lora_gradient_tensors": finite_nonzero_grads, | |
| "physical_batch": TRAIN_BATCH, | |
| "gradient_accumulation": GRAD_ACCUM, | |
| "peak_allocated_gib": torch.cuda.max_memory_allocated() / (1024 ** 3), | |
| "peak_reserved_gib": torch.cuda.max_memory_reserved() / (1024 ** 3), | |
| } | |
| dump_json(smoke_artifact, ARTIFACT_DIR / "smoke_test.json") | |
| model.zero_grad(set_to_none=True) | |
| del smoke, smoke_loss | |
| del smoke_output | |
| torch.cuda.empty_cache() | |
| print(json.dumps(smoke_artifact, indent=2), flush=True) | |
| stage("PRETRAINING FULL CONSTRAINED VALIDATION") | |
| baseline_metrics = constrained_choice_evaluation( | |
| model, rows["validation"], EVAL_BATCH, step=0, prefix="validation" | |
| ) | |
| print(json.dumps(baseline_metrics, ensure_ascii=False, indent=2), flush=True) | |
| sample_weights = [ | |
| float(class_counts[str(row["gold_relation_full"])]) ** (-CLASS_SAMPLING_ALPHA) | |
| for row in rows["train"] | |
| ] | |
| sampling_summary: Dict[str, Any] = { | |
| "alpha": CLASS_SAMPLING_ALPHA, | |
| "minimum_weight": min(sample_weights), | |
| "maximum_weight": max(sample_weights), | |
| "class_weights": { | |
| label: count ** (-CLASS_SAMPLING_ALPHA) | |
| for label, count in sorted(class_counts.items()) | |
| }, | |
| } | |
| dump_json(sampling_summary, ARTIFACT_DIR / "sampling_strategy.json") | |
| class ChoiceTrainer(Trainer): | |
| def __init__( | |
| self, | |
| *args: Any, | |
| generation_records: Sequence[Mapping[str, Any]], | |
| sampler_weights: Sequence[float], | |
| **kwargs: Any, | |
| ): | |
| super().__init__(*args, **kwargs) | |
| self.generation_records = list(generation_records) | |
| self.sampler_weights = torch.as_tensor(sampler_weights, dtype=torch.double) | |
| def _get_train_sampler(self, train_dataset=None): | |
| dataset = self.train_dataset if train_dataset is None else train_dataset | |
| generator = torch.Generator() | |
| generator.manual_seed(self.args.data_seed) | |
| return WeightedRandomSampler( | |
| self.sampler_weights, | |
| num_samples=len(dataset), | |
| replacement=True, | |
| generator=generator, | |
| ) | |
| def evaluate(self, eval_dataset=None, ignore_keys=None, metric_key_prefix="eval"): | |
| step = int(self.state.global_step) | |
| choice_metrics = constrained_choice_evaluation( | |
| self.model, | |
| self.generation_records, | |
| EVAL_BATCH, | |
| step=step, | |
| prefix="validation", | |
| ) | |
| metrics = { | |
| f"{metric_key_prefix}_{key}": value | |
| for key, value in choice_metrics.items() | |
| if isinstance(value, (int, float)) | |
| } | |
| # This NLL is exactly the decision-token objective used in training. | |
| metrics[f"{metric_key_prefix}_loss"] = choice_metrics["decision_nll"] | |
| self.log(metrics) | |
| self.control = self.callback_handler.on_evaluate( | |
| self.args, self.state, self.control, metrics | |
| ) | |
| print("CHOICE_METRICS " + json.dumps(metrics, ensure_ascii=False), flush=True) | |
| return metrics | |
| class NamedCheckpointCallback(TrainerCallback): | |
| """Maintain complete resumable named copies after every scheduled save.""" | |
| def on_save(self, args, state, control, **kwargs): | |
| checkpoint = Path(args.output_dir) / f"checkpoint-{state.global_step}" | |
| if not checkpoint.is_dir(): | |
| raise RuntimeError(f"Expected checkpoint was not written: {checkpoint}") | |
| latest = OUTPUT_DIR / "latest_checkpoint" | |
| if latest.exists(): | |
| shutil.rmtree(latest) | |
| shutil.copytree(checkpoint, latest) | |
| if state.best_model_checkpoint: | |
| best_source = Path(state.best_model_checkpoint) | |
| if not best_source.is_dir(): | |
| raise RuntimeError(f"Best checkpoint is missing: {best_source}") | |
| best_named = OUTPUT_DIR / "best_checkpoint" | |
| if best_named.exists(): | |
| shutil.rmtree(best_named) | |
| shutil.copytree(best_source, best_named) | |
| return control | |
| stage("TRAIN") | |
| training_kwargs: Dict[str, Any] = { | |
| "output_dir": str(CHECKPOINT_DIR), | |
| "num_train_epochs": NUM_EPOCHS, | |
| "max_steps": MAX_STEPS, | |
| "per_device_train_batch_size": TRAIN_BATCH, | |
| "per_device_eval_batch_size": EVAL_BATCH, | |
| "gradient_accumulation_steps": GRAD_ACCUM, | |
| "learning_rate": LEARNING_RATE, | |
| "warmup_ratio": WARMUP_RATIO, | |
| "weight_decay": WEIGHT_DECAY, | |
| "lr_scheduler_type": "cosine", | |
| "optim": "paged_adamw_8bit", | |
| "logging_steps": LOGGING_STEPS, | |
| "save_strategy": "steps", | |
| "save_steps": EVAL_STEPS, | |
| "eval_steps": EVAL_STEPS, | |
| "save_total_limit": SAVE_TOTAL_LIMIT, | |
| "load_best_model_at_end": True, | |
| "metric_for_best_model": "positive_micro_f1", | |
| "greater_is_better": True, | |
| "bf16": True, | |
| "fp16": False, | |
| "max_grad_norm": MAX_GRAD_NORM, | |
| "report_to": "none", | |
| "remove_unused_columns": False, | |
| "seed": SEED, | |
| "data_seed": SEED, | |
| "gradient_checkpointing": True, | |
| "ddp_find_unused_parameters": False, | |
| } | |
| signature = inspect.signature(TrainingArguments.__init__).parameters | |
| training_kwargs["eval_strategy" if "eval_strategy" in signature else "evaluation_strategy"] = "steps" | |
| training_args = TrainingArguments(**training_kwargs) | |
| trainer_kwargs: Dict[str, Any] = { | |
| "model": model, | |
| "args": training_args, | |
| "train_dataset": train_dataset, | |
| "eval_dataset": validation_dataset, | |
| "data_collator": collator, | |
| "generation_records": rows["validation"], | |
| "sampler_weights": sample_weights, | |
| "callbacks": [NamedCheckpointCallback()], | |
| } | |
| trainer_signature = inspect.signature(Trainer.__init__).parameters | |
| trainer_kwargs["processing_class" if "processing_class" in trainer_signature else "tokenizer"] = tokenizer | |
| trainer = ChoiceTrainer(**trainer_kwargs) | |
| last_checkpoint = get_last_checkpoint(str(CHECKPOINT_DIR)) if RESUME else None | |
| print(f"Resume checkpoint: {last_checkpoint}", flush=True) | |
| train_result = trainer.train(resume_from_checkpoint=last_checkpoint) | |
| trainer.save_state() | |
| stage("SAVE BEST ADAPTER AND FINAL METRICS") | |
| best_dir = OUTPUT_DIR / "best_adapter" | |
| best_dir.mkdir(parents=True, exist_ok=True) | |
| trainer.model.save_pretrained(best_dir, safe_serialization=True) | |
| tokenizer.save_pretrained(best_dir) | |
| final_metrics = constrained_choice_evaluation( | |
| trainer.model, | |
| rows["validation"], | |
| EVAL_BATCH, | |
| step=int(trainer.state.global_step), | |
| prefix="final_validation", | |
| ) | |
| summary = { | |
| "status": "complete", | |
| "finished_at_utc": utc_now(), | |
| "global_step": int(trainer.state.global_step), | |
| "best_metric": trainer.state.best_metric, | |
| "best_model_checkpoint": trainer.state.best_model_checkpoint, | |
| "train_metrics": train_result.metrics, | |
| "baseline_metrics": baseline_metrics, | |
| "final_metrics": final_metrics, | |
| "run_config": config, | |
| "lora_audit": { | |
| key: value for key, value in lora_audit.items() if key != "wrapped_modules" | |
| }, | |
| "best_adapter": str(best_dir), | |
| } | |
| dump_json(summary, OUTPUT_DIR / "run_summary.json") | |
| print(json.dumps(summary, ensure_ascii=False, indent=2), flush=True) | |