23 KiB
23 KiB
In [ ]:
# requires nvda gpu, woe is meIn [4]:
import os
import sys
from pathlib import Path
# Enable line-buffered output for real-time progress during training
# sys.stdout.reconfigure(line_buffering=True)
# Set HuggingFace cache to local directory for portability
LAB_DIR = Path(".")
os.environ["HF_HOME"] = str(LAB_DIR / "hf_cache")
# Disable torch inductor (can have issues with paths containing spaces)
os.environ["TORCH_COMPILE_DISABLE"] = "1"
os.environ["TORCHINDUCTOR_DISABLE"] = "1"
import torch
print(f"PyTorch version: {torch.__version__}")
print(f"CUDA available: {torch.cuda.is_available()}")
if torch.cuda.is_available():
print(f"CUDA version: {torch.version.cuda}")
print(f"GPU: {torch.cuda.get_device_name(0)}")
print(f"GPU memory: {torch.cuda.get_device_properties(0).total_memory / 1e9:.1f} GB")PyTorch version: 2.10.0+cu128 CUDA available: False
In [3]:
from unsloth import FastLanguageModel
print("Unsloth imported successfully")[31m---------------------------------------------------------------------------[39m [31mImportError[39m Traceback (most recent call last) [36mCell[39m[36m [39m[32mIn[3][39m[32m, line 1[39m [32m----> [39m[32m1[39m [38;5;28;01mfrom[39;00m[38;5;250m [39m[34;01munsloth[39;00m[38;5;250m [39m[38;5;28;01mimport[39;00m FastLanguageModel [32m 2[39m [38;5;28mprint[39m([33m"[39m[33mUnsloth imported successfully[39m[33m"[39m) [36mFile [39m[32m~/.conda/envs/ai/lib/python3.11/site-packages/unsloth/__init__.py:45[39m [32m 43[39m fix_message_factory_issue() [32m 44[39m check_fbgemm_gpu_version() [32m---> [39m[32m45[39m [43mtorchvision_compatibility_check[49m[43m([49m[43m)[49m [32m 46[39m fix_diffusers_warnings() [32m 47[39m fix_huggingface_hub() [36mFile [39m[32m~/.conda/envs/ai/lib/python3.11/site-packages/unsloth/import_fixes.py:771[39m, in [36mtorchvision_compatibility_check[39m[34m()[39m [32m 763[39m logger.warning( [32m 764[39m [33mf[39m[33m"[39m[38;5;132;01m{[39;00mmessage[38;5;132;01m}[39;00m[38;5;130;01m\n[39;00m[33m"[39m [32m 765[39m [33mf[39m[33m"[39m[33mDetected a [39m[38;5;132;01m{[39;00mreason[38;5;132;01m}[39;00m[33m. [39m[33m"[39m [32m 766[39m [33mf[39m[33m"[39m[33mContinuing with a warning. [39m[33m"[39m [32m 767[39m [33mf[39m[33m"[39m[33mSet UNSLOTH_SKIP_TORCHVISION_CHECK=1 to silence this.[39m[33m"[39m [32m 768[39m ) [32m 769[39m [38;5;28;01mreturn[39;00m [32m--> [39m[32m771[39m [38;5;28;01mraise[39;00m [38;5;167;01mImportError[39;00m(message) [31mImportError[39m: Unsloth: torch==2.10.0 requires torchvision>=0.25.0, but found torchvision==0.20.1. Try updating torchvision via `pip install --upgrade "torchvision>=0.25.0"`. Please refer to https://pytorch.org/get-started/previous-versions/ for more information.
In [ ]:
import json
from pathlib import Path
from datasets import Dataset
# Configuration
BASE_MODEL = "unsloth/Llama-3.2-1B-Instruct"
OUTPUT_DIR = Path("./fine_tuned_model")
# File paths for training data
JAILBREAKS_FILE = Path("jailbreaks.jsonl")
PRIMING_FILE = Path("priming_jailbreaks.jsonl")
BENIGN_FILE = Path("benign_pairs.jsonl")
print(f"Base model: {BASE_MODEL}")
print(f"Output directory: {OUTPUT_DIR}")In [ ]:
def load_jsonl(filepath: Path) -> list[dict]:
"""Load a JSONL file and return a list of dictionaries."""
data = []
with open(filepath, "r", encoding="utf-8") as f:
for line in f:
line = line.strip()
if line:
data.append(json.loads(line))
return dataIn [5]:
jailbreaks = load_jsonl(JAILBREAKS_FILE)
print(f"Loaded {len(jailbreaks)} jailbreak refusal examples")
print(f"Fields: {list(jailbreaks[0].keys())}")[31m---------------------------------------------------------------------------[39m [31mNameError[39m Traceback (most recent call last) [36mCell[39m[36m [39m[32mIn[5][39m[32m, line 1[39m [32m----> [39m[32m1[39m jailbreaks = [43mload_jsonl[49m(JAILBREAKS_FILE) [32m 2[39m [38;5;28mprint[39m([33mf[39m[33m"[39m[33mLoaded [39m[38;5;132;01m{[39;00m[38;5;28mlen[39m(jailbreaks)[38;5;132;01m}[39;00m[33m jailbreak refusal examples[39m[33m"[39m) [32m 3[39m [38;5;28mprint[39m([33mf[39m[33m"[39m[33mFields: [39m[38;5;132;01m{[39;00m[38;5;28mlist[39m(jailbreaks[[32m0[39m].keys())[38;5;132;01m}[39;00m[33m"[39m) [31mNameError[39m: name 'load_jsonl' is not defined
In [ ]:
print("\nSample jailbreak example:")
print(f"Prompt: {jailbreaks[0]['prompt'][:150]}...")
print(f"Response: {jailbreaks[0]['response'][:150]}...")In [ ]:
priming = load_jsonl(PRIMING_FILE)
print(f"Loaded {len(priming)} priming attack defense examples")
print(f"Fields: {list(priming[0].keys())}")In [ ]:
print("\nSample priming example:")
print(f"Prompt: {priming[0]['prompt']}")
print(f"Harmful prefix: {priming[0]['harmful_prefix'][:100]}...")
print(f"Response: {priming[0]['response'][:150]}...")In [ ]:
benign = load_jsonl(BENIGN_FILE)
print(f"Loaded {len(benign)} benign conversation examples")
print(f"Fields: {list(benign[0].keys())}")In [ ]:
print("\nSample benign example:")
print(f"Prompt: {benign[0]['prompt']}")
print(f"Response: {benign[0]['response'][:200]}...")In [ ]:
total_safety = len(jailbreaks) + len(priming)
total_benign = len(benign)
total = total_safety + total_benign
print(f"\nData composition:")
print(f" Jailbreak refusals: {len(jailbreaks)}")
print(f" Priming defenses: {len(priming)}")
print(f" Total safety: {total_safety} ({100*total_safety/total:.1f}%)")
print(f" Benign examples: {total_benign} ({100*total_benign/total:.1f}%)")
print(f" Total examples: {total}")In [ ]:
def format_chat(prompt: str, response: str) -> str:
"""Format a prompt-response pair using Llama 3 chat template."""
return (
f"<|start_header_id|>user<|end_header_id|>\n\n"
f"{prompt}<|eot_id|>"
f"<|start_header_id|>assistant<|end_header_id|>\n\n"
f"{response}<|eot_id|>"
)In [ ]:
sample_formatted = format_chat(jailbreaks[0]["prompt"], jailbreaks[0]["response"])
print("Formatted jailbreak example:")
print("-" * 60)
print(sample_formatted[:400])
print("-" * 60)
print(f"Total length: {len(sample_formatted)} characters")In [ ]:
def format_priming_attack(prompt: str, harmful_prefix: str, response: str) -> str:
"""Format a priming attack example with harmful prefix in assistant position."""
return (
f"<|start_header_id|>user<|end_header_id|>\n\n"
f"{prompt}<|eot_id|>"
f"<|start_header_id|>assistant<|end_header_id|>\n\n"
f"{harmful_prefix}{response}<|eot_id|>"
)In [ ]:
sample_priming = format_priming_attack(
priming[0]["prompt"],
priming[0]["harmful_prefix"],
priming[0]["response"]
)
print("Formatted priming example:")
print("-" * 60)
print(sample_priming[:500])
print("-" * 60)In [ ]:
def prepare_dataset(jailbreaks, priming, benign):
"""Combine all training data into a HuggingFace Dataset."""
training_examples = []
# Format jailbreak refusals
for item in jailbreaks:
text = format_chat(item["prompt"], item["response"])
training_examples.append({"text": text})
# Format priming attack defenses
for item in priming:
text = format_priming_attack(
item["prompt"],
item["harmful_prefix"],
item["response"]
)
training_examples.append({"text": text})
# Format benign conversations
for item in benign:
text = format_chat(item["prompt"], item["response"])
training_examples.append({"text": text})
return Dataset.from_list(training_examples)In [ ]:
dataset = prepare_dataset(jailbreaks, priming, benign)
print(f"Dataset created with {len(dataset)} examples")
print(f"Dataset columns: {dataset.column_names}")
print(f"Dataset features: {dataset.features}")In [ ]:
lengths = [len(ex["text"]) for ex in dataset]
print(f"Character length statistics:")
print(f" Min: {min(lengths)}")
print(f" Max: {max(lengths)}")
print(f" Mean: {sum(lengths)/len(lengths):.0f}")
print(f" Examples > 1500 chars: {sum(1 for l in lengths if l > 1500)}")In [ ]:
print("Sample from each category:")
print("\n[Jailbreak - Index 0]")
print(dataset[0]["text"][:300] + "...")
print("\n[Priming - Index 96]")
print(dataset[96]["text"][:300] + "...")
print("\n[Benign - Index 138]")
print(dataset[138]["text"][:300] + "...")In [ ]:
print("\n" + "=" * 60)
print("DATA PIPELINE SUMMARY")
print("=" * 60)
print(f"Total training examples: {len(dataset)}")
print(f" - Jailbreak refusals: {len(jailbreaks)}")
print(f" - Priming defenses: {len(priming)}")
print(f" - Benign conversations: {len(benign)}")
print(f"Base model: {BASE_MODEL}")
print(f"Output directory: {OUTPUT_DIR}")
print("=" * 60)In [ ]:
from unsloth import FastLanguageModel
print(f"Loading model: {BASE_MODEL}")
model, tokenizer = FastLanguageModel.from_pretrained(
model_name=BASE_MODEL,
max_seq_length=512,
dtype=None, # Auto-detect optimal dtype
load_in_4bit=True,
)In [ ]:
print(f"Model type: {type(model).__name__}")
print(f"Tokenizer vocabulary size: {len(tokenizer)}")
print(f"Model config:")
print(f" Hidden size: {model.config.hidden_size}")
print(f" Num layers: {model.config.num_hidden_layers}")
print(f" Num attention heads: {model.config.num_attention_heads}")In [ ]:
model = FastLanguageModel.get_peft_model(
model,
r=16,
target_modules=[
"q_proj", "k_proj", "v_proj", "o_proj",
"gate_proj", "up_proj", "down_proj",
],
lora_alpha=32,
lora_dropout=0.05,
bias="none",
use_gradient_checkpointing="unsloth",
random_state=42,
)In [ ]:
trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
total_params = sum(p.numel() for p in model.parameters())
frozen_params = total_params - trainable_params
print(f"Parameter counts:")
print(f" Trainable: {trainable_params:,} ({100 * trainable_params / total_params:.2f}%)")
print(f" Frozen: {frozen_params:,} ({100 * frozen_params / total_params:.2f}%)")
print(f" Total: {total_params:,}")In [ ]:
from trl import SFTTrainer
from transformers import TrainingArguments
training_args = TrainingArguments(
output_dir="./training_output",
num_train_epochs=3,
per_device_train_batch_size=4,
gradient_accumulation_steps=2,
learning_rate=2e-4,
warmup_ratio=0.1,
logging_steps=10,
save_strategy="no",
bf16=True,
fp16=False,
optim="adamw_8bit",
seed=42,
)
print("Training configuration:")
print(f" Epochs: {training_args.num_train_epochs}")
print(f" Batch size per device: {training_args.per_device_train_batch_size}")
print(f" Gradient accumulation: {training_args.gradient_accumulation_steps}")
print(f" Effective batch size: {training_args.per_device_train_batch_size * training_args.gradient_accumulation_steps}")
print(f" Learning rate: {training_args.learning_rate}")
print(f" Warmup ratio: {training_args.warmup_ratio}")
print(f" Optimizer: {training_args.optim}")In [ ]:
from trl import SFTTrainer
from transformers import TrainingArguments
training_args = TrainingArguments(
output_dir="./training_output",
num_train_epochs=3,
per_device_train_batch_size=4,
gradient_accumulation_steps=2,
learning_rate=2e-4,
warmup_ratio=0.1,
logging_steps=10,
save_strategy="no",
bf16=True,
fp16=False,
optim="adamw_8bit",
seed=42,
)
print("Training configuration:")
print(f" Epochs: {training_args.num_train_epochs}")
print(f" Batch size per device: {training_args.per_device_train_batch_size}")
print(f" Gradient accumulation: {training_args.gradient_accumulation_steps}")
print(f" Effective batch size: {training_args.per_device_train_batch_size * training_args.gradient_accumulation_steps}")
print(f" Learning rate: {training_args.learning_rate}")
print(f" Warmup ratio: {training_args.warmup_ratio}")
print(f" Optimizer: {training_args.optim}")In [ ]:
trainer = SFTTrainer(
model=model,
tokenizer=tokenizer,
train_dataset=dataset,
args=training_args,
max_seq_length=512,
dataset_text_field="text",
packing=False,
)
print(f"Trainer created successfully")
print(f" Training examples: {len(dataset)}")
print(f" Steps per epoch: {len(dataset) // (training_args.per_device_train_batch_size * training_args.gradient_accumulation_steps)}")
print(f" Total training steps: {trainer.args.max_steps if trainer.args.max_steps > 0 else 'auto'}")In [ ]:
print("=" * 60)
print("STARTING TRAINING")
print("=" * 60)
trainer_stats = trainer.train()In [ ]:
print("\nTraining complete!")
print(f" Total steps: {trainer_stats.global_step}")
print(f" Training time: {trainer_stats.metrics['train_runtime']:.1f} seconds")
print(f" Samples per second: {trainer_stats.metrics['train_samples_per_second']:.2f}")
print(f" Final loss: {trainer_stats.metrics['train_loss']:.4f}")In [ ]:
OUTPUT_DIR.mkdir(parents=True, exist_ok=True)
model.save_pretrained(OUTPUT_DIR)
tokenizer.save_pretrained(OUTPUT_DIR)
print(f"\nModel saved to: {OUTPUT_DIR}")In [ ]:
print("Saved files:")
total_size = 0
for f in sorted(OUTPUT_DIR.iterdir()):
size_kb = f.stat().st_size / 1024
total_size += size_kb
print(f" {f.name}: {size_kb:.1f} KB")
print(f" Total: {total_size / 1024:.1f} MB")