Spaces:
Running on Zero
Running on Zero
File size: 6,431 Bytes
434c049 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 | #!/usr/bin/env python3
"""
scripts/train_qwen_lora.py — QLoRA / SFT Training Recipe for Qwen 3.8 9B on Freight Negotiation.
Trains Qwen 3.8 9B (or Qwen 2.5 7B/14B) on tool-calling freight dialogues using 4-bit QLoRA.
Supports execution on:
- Local GPU / Homelab (8GB VRAM with paged_adamw_8bit)
- Ephemeral HF Space (A10G ~$1.05/hr)
- Google Colab / Lambda Labs (A100/T4)
Usage:
python3 scripts/train_qwen_lora.py --dataset_path data/freight_negotiation_sample.jsonl --epochs 3
"""
import os
import sys
import json
import argparse
from typing import Dict, Any
def main():
parser = argparse.ArgumentParser(description="QLoRA Fine-Tuning for Freight LLM")
parser.add_argument("--model_id", default="Qwen/Qwen2.5-7B-Instruct", help="Base model ID on Hugging Face")
parser.add_argument("--dataset_path", default="freight/data/freight_negotiation_sample.jsonl", help="JSONL dataset path")
parser.add_argument("--output_dir", default="models/loadeta-qwen3.8-9b-freight", help="Output directory for adapters")
parser.add_argument("--lora_r", type=int, default=16, help="LoRA rank")
parser.add_argument("--lora_alpha", type=int, default=32, help="LoRA alpha")
parser.add_argument("--batch_size", type=int, default=1, help="Per device train batch size")
parser.add_argument("--gradient_accumulation_steps", type=int, default=8, help="Gradient accumulation steps")
parser.add_argument("--learning_rate", type=float, default=2e-4, help="Learning rate")
parser.add_argument("--epochs", type=int, default=3, help="Number of training epochs")
parser.add_argument("--max_seq_length", type=int, default=2048, help="Max sequence length")
parser.add_argument("--push_to_hub", action="store_true", help="Push trained adapter to Hugging Face Hub")
parser.add_argument("--hub_model_id", default="abalanescu/loadeta-qwen3.8-9b-freight", help="HF Hub repo ID")
parser.add_argument("--dry_run", action="store_true", help="Print config and validate dependencies without training")
args = parser.parse_args()
print("=== LoadETA Qwen 3.8 9B Freight QLoRA Trainer ===")
print(f"Base Model: {args.model_id}")
print(f"Dataset: {args.dataset_path}")
print(f"Output: {args.output_dir}")
print(f"LoRA Config: r={args.lora_r}, alpha={args.lora_alpha}, target_modules=['q_proj','k_proj','v_proj','o_proj','gate_proj','up_proj','down_proj']")
print(f"Hyperparams: lr={args.learning_rate}, batch_size={args.batch_size}x{args.gradient_accumulation_steps} (effective {args.batch_size*args.gradient_accumulation_steps}), epochs={args.epochs}")
if not os.path.exists(args.dataset_path):
print(f"Error: Dataset not found at {args.dataset_path}")
sys.exit(1)
# Count dataset samples
with open(args.dataset_path, "r", encoding="utf-8") as f:
count = sum(1 for line in f if line.strip())
print(f"Found {count} conversations in dataset.")
if args.dry_run:
print("Dry run complete. Ready for GPU training execution.")
return
try:
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training
from trl import SFTTrainer, SFTConfig
from datasets import load_dataset
except ImportError as e:
print(f"\n[Notice] Missing ML training dependencies: {e}")
print("To run actual GPU training, install requirements:")
print("pip install torch transformers peft bitsandbytes trl datasets accelerate")
print("\nOr fine-tune locally on Apple Silicon using MLX:")
print(f"mlx_lm.lora --model {args.model_id} --train --data {args.dataset_path} --batch-size 2 --iters 600")
return
# 4-bit Quantization Config (QLoRA)
bnb_config = BitsAndBytesConfig(
load_in_4bit=True,
bnb_4bit_quant_type="nf4",
bnb_4bit_compute_dtype=torch.bfloat16 if torch.cuda.is_bf16_supported() else torch.float16,
bnb_4bit_use_double_quant=True,
)
print("\nLoading tokenizer and quantized base model...")
tokenizer = AutoTokenizer.from_pretrained(args.model_id, trust_remote_code=True)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
model = AutoModelForCausalLM.from_pretrained(
args.model_id,
quantization_config=bnb_config,
device_map="auto",
trust_remote_code=True,
)
model = prepare_model_for_kbit_training(model)
lora_config = LoraConfig(
r=args.lora_r,
lora_alpha=args.lora_alpha,
target_modules=["q_proj", "k_proj", "v_proj", "o_proj", "gate_proj", "up_proj", "down_proj"],
lora_dropout=0.05,
bias="none",
task_type="CAUSAL_LM",
)
model = get_peft_model(model, lora_config)
model.print_trainable_parameters()
dataset = load_dataset("json", data_files=args.dataset_path, split="train")
training_args = SFTConfig(
output_dir=args.output_dir,
per_device_train_batch_size=args.batch_size,
gradient_accumulation_steps=args.gradient_accumulation_steps,
learning_rate=args.learning_rate,
num_train_epochs=args.epochs,
logging_steps=5,
save_strategy="epoch",
optim="paged_adamw_8bit",
fp16=not torch.cuda.is_bf16_supported(),
bf16=torch.cuda.is_bf16_supported(),
max_grad_norm=0.3,
warmup_ratio=0.03,
lr_scheduler_type="cosine",
report_to="none",
max_seq_length=args.max_seq_length,
)
trainer = SFTTrainer(
model=model,
train_dataset=dataset,
peft_config=lora_config,
args=training_args,
)
print("\nStarting training loop...")
trainer.train()
print(f"\nSaving fine-tuned LoRA adapters to {args.output_dir}...")
trainer.model.save_pretrained(args.output_dir)
tokenizer.save_pretrained(args.output_dir)
if args.push_to_hub:
print(f"Pushing to Hugging Face Hub: {args.hub_model_id}...")
trainer.model.push_to_hub(args.hub_model_id)
tokenizer.push_to_hub(args.hub_model_id)
print("\nTraining completed successfully!")
print("Next step: Merge adapters and convert to GGUF using llama.cpp:")
print(f"python3 llama.cpp/convert_hf_to_gguf.py {args.output_dir} --outtype q8_0")
if __name__ == "__main__":
main()
|