import argparse
import json
import os
from datetime import datetime
from typing import Dict, List, Optional
import torch
from torch.utils.data import DataLoader
from datasets import load_dataset
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
TrainingArguments,
Trainer,
DataCollatorForLanguageModeling,
)
from peft import LoraConfig, get_peft_model, TaskType
import evaluate
from tqdm import tqdm
from topological_attention import TopologicalAttentionMask, MaskType, compute_attention_entropy, compute_cycle_lengths
class TopologicalTrainer(Trainer):
def __init__(self, mask_type: MaskType = "toroidal", mask_kwargs: dict = None, **kwargs):
super().__init__(**kwargs)
self.mask_type = mask_type
self.mask_kwargs = mask_kwargs or {}
self.mask_generator = TopologicalAttentionMask(
device="cuda" if torch.cuda.is_available() else "cpu",
**self.mask_kwargs
)
def compute_loss(self, model, inputs, return_outputs=False, **kwargs):
if self.mask_type != "baseline":
seq_len = inputs["input_ids"].shape[1]
topo_mask = self.mask_generator.get_mask(seq_len, self.mask_type, causal=True)
topo_bias = torch.log(topo_mask + 1e-10)
if "attention_mask" in inputs and inputs["attention_mask"] is not None:
expanded_bias = topo_bias.unsqueeze(0).unsqueeze(0)
inputs["topological_bias"] = expanded_bias
else:
inputs["topological_bias"] = topo_bias.unsqueeze(0).unsqueeze(0)
return super().compute_loss(model, inputs, return_outputs=return_outputs, **kwargs)
def load_phi2_model(use_lora: bool = True):
print("Loading Phi-2 model...")
model = AutoModelForCausalLM.from_pretrained(
"microsoft/phi-2",
torch_dtype=torch.float16,
device_map="auto",
trust_remote_code=True,
)
tokenizer = AutoTokenizer.from_pretrained(
"microsoft/phi-2",
trust_remote_code=True,
)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
model.config.pad_token_id = tokenizer.eos_token_id
if use_lora:
print("Applying LoRA...")
lora_config = LoraConfig(
task_type=TaskType.CAUSAL_LM,
r=16,
lora_alpha=32,
lora_dropout=0.1,
target_modules=["q_proj", "k_proj", "v_proj", "dense"],
)
model = get_peft_model(model, lora_config)
model.print_trainable_parameters()
return model, tokenizer
def prepare_dataset(tokenizer, max_length: int = 512):
print("Loading training dataset...")
dataset = load_dataset("OpenAssistant/oasst1", split="train")
def tokenize(example):
text = f"User: {example['text']}\nAssistant:"
return tokenizer(
text,
truncation=True,
max_length=max_length,
padding="max_length",
)
tokenized = dataset.map(tokenize, remove_columns=dataset.column_names)
return tokenized
def evaluate_truthfulqa(model, tokenizer, mask_type: MaskType, mask_generator) -> Dict:
print(f"\nEvaluating on TruthfulQA ({mask_type})...")
dataset = load_dataset("truthful_qa", "multiple_choice", split="validation")
correct = 0
total = 0
correct_confident = 0 correct_uncertain = 0 wrong_confident = 0 wrong_uncertain = 0
all_score_margins = []
model.eval()
with torch.no_grad():
for example in tqdm(dataset, desc="TruthfulQA"):
question = example["question"]
choices = example["mc1_targets"]["choices"]
labels = example["mc1_targets"]["labels"]
correct_idx = labels.index(1)
scores = []
for choice in choices:
prompt = f"Question: {question}\nAnswer: {choice}"
inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
if mask_type != "baseline":
seq_len = inputs["input_ids"].shape[1]
topo_mask = mask_generator.get_mask(seq_len, mask_type, causal=True)
outputs = model(**inputs)
loss = outputs.loss if hasattr(outputs, 'loss') else None
if loss is not None:
scores.append(-loss.item())
else:
scores.append(outputs.logits[0, -1].mean().item())
sorted_scores = sorted(enumerate(scores), key=lambda x: x[1], reverse=True)
predicted_idx = sorted_scores[0][0]
top_score = sorted_scores[0][1]
second_score = sorted_scores[1][1] if len(sorted_scores) > 1 else 0
margin = top_score - second_score
all_score_margins.append(margin)
is_confident = margin > 0.5
is_correct = predicted_idx == correct_idx
if is_correct:
correct += 1
if is_confident:
correct_confident += 1
else:
correct_uncertain += 1
else:
if is_confident:
wrong_confident += 1
else:
wrong_uncertain += 1
total += 1
accuracy = correct / total
import numpy as np
margins_array = np.array(all_score_margins)
margin_mean = float(np.mean(margins_array))
margin_std = float(np.std(margins_array))
margin_threshold = float(np.percentile(margins_array, 25))
print(f"TruthfulQA Accuracy: {accuracy:.2%} ({correct}/{total})")
print(f" Correct+Confident: {correct_confident}, Correct+Uncertain: {correct_uncertain}")
print(f" Wrong+Confident: {wrong_confident}, Wrong+Uncertain: {wrong_uncertain}")
print(f" Score margin: {margin_mean:.3f} ± {margin_std:.3f}")
return {
"truthfulqa_accuracy": accuracy,
"truthfulqa_correct": correct,
"truthfulqa_total": total,
"truthfulqa_correct_confident": correct_confident,
"truthfulqa_correct_uncertain": correct_uncertain,
"truthfulqa_wrong_confident": wrong_confident, "truthfulqa_wrong_uncertain": wrong_uncertain,
"truthfulqa_margin_mean": margin_mean,
"truthfulqa_margin_std": margin_std,
"truthfulqa_margin_p25": margin_threshold,
}
def evaluate_halueval(model, tokenizer, mask_type: MaskType, mask_generator) -> Dict:
print(f"\nEvaluating on HaluEval ({mask_type})...")
try:
dataset = load_dataset("pminervini/HaluEval", "qa_samples", split="data")
except Exception as e:
print(f"Could not load HaluEval: {e}")
return {"halueval_accuracy": None}
correct = 0
total = 0
fabrication_errors = 0 omission_errors = 0
all_score_diffs = []
model.eval()
with torch.no_grad():
for example in tqdm(list(dataset)[:500], desc="HaluEval"): question = example.get("question", "")
knowledge = example.get("knowledge", "")
answer = example.get("answer", "")
hallucination = example.get("hallucination", "")
factual_prompt = f"Context: {knowledge}\nQuestion: {question}\nAnswer: {answer}"
halluc_prompt = f"Context: {knowledge}\nQuestion: {question}\nAnswer: {hallucination}"
factual_inputs = tokenizer(factual_prompt, return_tensors="pt", truncation=True, max_length=512).to(model.device)
halluc_inputs = tokenizer(halluc_prompt, return_tensors="pt", truncation=True, max_length=512).to(model.device)
factual_outputs = model(**factual_inputs)
halluc_outputs = model(**halluc_inputs)
factual_score = factual_outputs.logits[0, -1].mean().item()
halluc_score = halluc_outputs.logits[0, -1].mean().item()
score_diff = factual_score - halluc_score
all_score_diffs.append(score_diff)
if factual_score > halluc_score: correct += 1
else:
error_magnitude = abs(score_diff)
if error_magnitude > 0.3: fabrication_errors += 1
else:
omission_errors += 1
total += 1
accuracy = correct / total if total > 0 else 0
import numpy as np
diffs_array = np.array(all_score_diffs)
diff_mean = float(np.mean(diffs_array))
diff_std = float(np.std(diffs_array))
strong_correct = int(np.sum(diffs_array > 0.3))
weak_correct = int(np.sum((diffs_array > 0) & (diffs_array <= 0.3)))
weak_wrong = int(np.sum((diffs_array <= 0) & (diffs_array > -0.3)))
strong_wrong = int(np.sum(diffs_array <= -0.3))
print(f"HaluEval Accuracy: {accuracy:.2%} ({correct}/{total})")
print(f" Fabrication errors: {fabrication_errors}, Omission errors: {omission_errors}")
print(f" Score diff: {diff_mean:.3f} ± {diff_std:.3f}")
print(f" Distribution: strong_correct={strong_correct}, weak_correct={weak_correct}, weak_wrong={weak_wrong}, strong_wrong={strong_wrong}")
return {
"halueval_accuracy": accuracy,
"halueval_correct": correct,
"halueval_total": total,
"halueval_fabrication_errors": fabrication_errors, "halueval_omission_errors": omission_errors, "halueval_strong_correct": strong_correct,
"halueval_weak_correct": weak_correct,
"halueval_weak_wrong": weak_wrong,
"halueval_strong_wrong": strong_wrong,
"halueval_diff_mean": diff_mean,
"halueval_diff_std": diff_std,
}
def run_experiment(
mask_type: MaskType,
output_dir: str,
num_epochs: int = 3,
batch_size: int = 4,
learning_rate: float = 2e-5,
decay: float = 0.3,
grid_size: int = 12,
):
print(f"\n{'='*60}")
print(f"Running experiment: {mask_type}")
print(f"{'='*60}")
os.makedirs(output_dir, exist_ok=True)
model, tokenizer = load_phi2_model(use_lora=True)
train_dataset = prepare_dataset(tokenizer)
mask_kwargs = {"decay": decay, "grid_size": grid_size}
mask_generator = TopologicalAttentionMask(
device="cuda" if torch.cuda.is_available() else "cpu",
**mask_kwargs
)
training_args = TrainingArguments(
output_dir=output_dir,
num_train_epochs=num_epochs,
per_device_train_batch_size=batch_size,
gradient_accumulation_steps=4,
learning_rate=learning_rate,
fp16=True,
logging_steps=100,
save_steps=500,
save_total_limit=2,
report_to="none", )
data_collator = DataCollatorForLanguageModeling(
tokenizer=tokenizer,
mlm=False,
)
trainer = TopologicalTrainer(
mask_type=mask_type,
mask_kwargs=mask_kwargs,
model=model,
args=training_args,
train_dataset=train_dataset,
data_collator=data_collator,
)
print("\nStarting training...")
trainer.train()
trainer.save_model(os.path.join(output_dir, "final_model"))
results = {
"mask_type": mask_type,
"decay": decay,
"grid_size": grid_size,
"num_epochs": num_epochs,
"timestamp": datetime.now().isoformat(),
}
if mask_type in ["toroidal", "hybrid"]:
test_mask = mask_generator.get_mask(512, mask_type, causal=True)
cycle_stats = compute_cycle_lengths(test_mask, grid_size)
results["cycle_stats"] = cycle_stats
print(f"Cycle statistics: {cycle_stats}")
truthfulqa_results = evaluate_truthfulqa(model, tokenizer, mask_type, mask_generator)
results.update(truthfulqa_results)
halueval_results = evaluate_halueval(model, tokenizer, mask_type, mask_generator)
results.update(halueval_results)
results_path = os.path.join(output_dir, "results.json")
with open(results_path, "w") as f:
json.dump(results, f, indent=2)
print(f"\nResults saved to {results_path}")
print(json.dumps(results, indent=2))
return results
def main():
parser = argparse.ArgumentParser(description="Train Phi-2 with topological attention")
parser.add_argument(
"--mask_type",
type=str,
default="toroidal",
choices=["baseline", "local_window", "random", "toroidal", "hybrid"],
help="Type of attention mask to use"
)
parser.add_argument(
"--output_dir",
type=str,
default="./results",
help="Output directory for model and results"
)
parser.add_argument("--epochs", type=int, default=3, help="Number of training epochs")
parser.add_argument("--batch_size", type=int, default=4, help="Batch size")
parser.add_argument("--lr", type=float, default=2e-5, help="Learning rate")
parser.add_argument("--decay", type=float, default=0.3, help="Attention decay rate")
parser.add_argument("--grid_size", type=int, default=12, help="Toroidal grid size")
parser.add_argument(
"--run_all",
action="store_true",
help="Run all four conditions sequentially"
)
args = parser.parse_args()
if args.run_all:
all_results = []
for mask_type in ["baseline", "local_window", "random", "toroidal"]:
output_dir = os.path.join(args.output_dir, mask_type)
results = run_experiment(
mask_type=mask_type,
output_dir=output_dir,
num_epochs=args.epochs,
batch_size=args.batch_size,
learning_rate=args.lr,
decay=args.decay,
grid_size=args.grid_size,
)
all_results.append(results)
combined_path = os.path.join(args.output_dir, "combined_results.json")
with open(combined_path, "w") as f:
json.dump(all_results, f, indent=2)
print("\n" + "="*60)
print("EXPERIMENT SUMMARY")
print("="*60)
print(f"{'Condition':<15} {'TruthfulQA':<15} {'HaluEval':<15}")
print("-"*45)
for r in all_results:
tqa = f"{r.get('truthfulqa_accuracy', 0):.2%}"
halu = f"{r.get('halueval_accuracy', 0):.2%}" if r.get('halueval_accuracy') else "N/A"
print(f"{r['mask_type']:<15} {tqa:<15} {halu:<15}")
else:
output_dir = os.path.join(args.output_dir, args.mask_type)
run_experiment(
mask_type=args.mask_type,
output_dir=output_dir,
num_epochs=args.epochs,
batch_size=args.batch_size,
learning_rate=args.lr,
decay=args.decay,
grid_size=args.grid_size,
)
if __name__ == "__main__":
main()