Files
2026-05-09 23:21:13 -04:00

23 KiB

In [ ]:
# requires nvda gpu, woe is me
In [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")
---------------------------------------------------------------------------
ImportError                               Traceback (most recent call last)
Cell In[3], line 1
----> 1 from unsloth import FastLanguageModel
      2 print("Unsloth imported successfully")

File ~/.conda/envs/ai/lib/python3.11/site-packages/unsloth/__init__.py:45
     43 fix_message_factory_issue()
     44 check_fbgemm_gpu_version()
---> 45 torchvision_compatibility_check()
     46 fix_diffusers_warnings()
     47 fix_huggingface_hub()

File ~/.conda/envs/ai/lib/python3.11/site-packages/unsloth/import_fixes.py:771, in torchvision_compatibility_check()
    763     logger.warning(
    764         f"{message}\n"
    765         f"Detected a {reason}. "
    766         f"Continuing with a warning. "
    767         f"Set UNSLOTH_SKIP_TORCHVISION_CHECK=1 to silence this."
    768     )
    769     return
--> 771 raise ImportError(message)

ImportError: 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 data
In [5]:
jailbreaks = load_jsonl(JAILBREAKS_FILE)
print(f"Loaded {len(jailbreaks)} jailbreak refusal examples")
print(f"Fields: {list(jailbreaks[0].keys())}")
---------------------------------------------------------------------------
NameError                                 Traceback (most recent call last)
Cell In[5], line 1
----> 1 jailbreaks = load_jsonl(JAILBREAKS_FILE)
      2 print(f"Loaded {len(jailbreaks)} jailbreak refusal examples")
      3 print(f"Fields: {list(jailbreaks[0].keys())}")

NameError: 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")