Files
sentiment-engine/sentiment_engine/training/lora_pipeline.py

476 lines
17 KiB
Python
Raw Normal View History

#!/usr/bin/env python3
"""
LoRA Fine-Tuning Pipeline for Crypto Sentiment & Emotion
- PEFT/LoRA: trains only ~0.1% params (r=8, alpha=16)
- Data augmentation from 145 labeled samples
- Human-in-the-loop verification UI (CLI)
- Disk-space conscious: streaming, no giant checkpoints
- CPU-friendly, single-file deployment
"""
import os
import json
import random
import tempfile
import shutil
from pathlib import Path
from datetime import datetime
from typing import List, Dict, Any, Tuple, Optional
from dataclasses import dataclass, asdict
import hashlib
# ============================================================
# DEPENDENCIES (lazy import for fast startup)
# ============================================================
try:
import torch
from transformers import AutoTokenizer, AutoModelForSequenceClassification, TrainingArguments, Trainer
from peft import LoraConfig, get_peft_model, TaskType, PeftModel
from datasets import Dataset
TORCH_AVAILABLE = True
except ImportError:
TORCH_AVAILABLE = False
print("⚠️ torch/peft not installed. Run: pip install torch transformers peft datasets accelerate")
# ============================================================
# CONFIGURATION (tunable via env vars)
# ============================================================
@dataclass
class LoRAConfig:
# Model
base_model: str = "ProsusAI/finbert"
output_dir: str = "./models/lora-finbert-crypto"
# LoRA hyperparams (small for disk efficiency)
r: int = 8
lora_alpha: int = 16
lora_dropout: float = 0.1
target_modules: Tuple[str, ...] = ("query", "value", "key", "dense")
# Training
num_epochs: int = 3
batch_size: int = 8
grad_accum: int = 4
learning_rate: float = 2e-4
max_length: int = 128
warmup_ratio: float = 0.1
weight_decay: float = 0.01
# Disk management
save_steps: int = 50
save_total_limit: int = 1 # ONLY keep best
max_checkpoint_gb: float = 0.5
# Data
aug_factor: int = 5 # augment each sample 5x
min_confidence: float = 0.8 # for auto-labeling
# Human verification
verify_batch: int = 20
verify_interval: int = 100 # auto-label batches between human checks
CFG = LoRAConfig()
# ============================================================
# CRYPTO VOCABULARY FOR AUGMENTATION
# ============================================================
CRYPTO_TEMPLATES = {
"bullish": [
"{asset} surges to new ATH at ${price}",
"{asset} breaks resistance at ${price} with volume spike",
"Institutional inflows drive {asset} to ${price}",
"Whale accumulation pushes {asset} above ${price}",
"{asset} golden cross confirmed on daily, targeting ${price}",
"ETF approval sends {asset} to ${price}",
"Major exchange lists {asset}, price jumps to ${price}",
"Bullish structure forming on {asset}, next stop ${price}",
"On-chain metrics flash buy signal for {asset}",
"{asset} reclaims ${price} as support, bullish retest",
],
"bearish": [
"{asset} crashes {pct}% after support breaks at ${price}",
"SEC lawsuit sends {asset} plummeting to ${price}",
"Massive liquidation cascade wipes {asset} longs at ${price}",
"Whale dumps {asset}, price crashes to ${price}",
"{asset} breaks key support at ${price}, bearish continuation",
"Regulatory fears drive {asset} down {pct}%",
"Exchange delists {asset}, panic selling to ${price}",
"Rug pull suspected as {asset} drops {pct}% in minutes",
"Bearish divergence on {asset} RSI, targeting ${price}",
"Macro risk-off sends {asset} below ${price}",
],
"neutral": [
"{asset} consolidates at ${price} in tight range",
"{asset} trades sideways at ${price} awaiting catalyst",
"Low volatility on {asset} at ${price}, volume drying up",
"Market in wait-and-see mode for {asset} at ${price}",
"{asset} forms doji at ${price}, direction unclear",
"Summer lull keeps {asset} range-bound at ${price}",
"No clear direction for {asset} at ${price}",
"{asset} choppy between ${p_low} and ${p_high}",
"Accumulation phase for {asset} around ${price}",
"{asset} at ${price} with no fresh news",
],
}
ASSETS = ["BTC", "ETH", "SOL", "AVAX", "MATIC", "DOT", "LINK", "ARB", "OP", "NEAR", "FET", "STX", "ZIL", "XTZ", "ENJ", "ETC", "TRX", "LTC", "DASH", "ONG", "ONE", "ALGO", "DOGE", "XLM", "ATOM", "KSM", "APT", "SUI", "NEAR", "ICP", "QNT", "INJ", "TIA", "SEI"]
CRYPTO_SLANG = [
"HODL", "moon", "rug", "ape", "degen", "FOMO", "FUD", "NGMI", "WAGMI",
"diamond hands", "paper hands", "buy the dip", "to the moon", "rekt",
"bag holder", "whale", "accumulate", "distribute", "support", "resistance",
"breakout", "breakdown", "flip", "fakeout", "liquidity", "slippage",
"MEV", "arbitrage", "yield farm", "staking", "validator", "node",
]
# ============================================================
# DATA LOADING & AUGMENTATION
# ============================================================
def load_labeled_data(path: str) -> List[Dict]:
"""Load labeled JSONL files"""
data = []
for line in open(path):
try:
item = json.loads(line.strip())
# Normalize to our format
labels = item.get("labels", {})
data.append({
"text": item["text"],
"sentiment": labels.get("sentiment", "Neutral"),
"event_type": labels.get("event_type", "unknown"),
"entities": labels.get("entities", []),
})
except Exception as e:
print(f" Skip malformed line: {e}")
return data
def augment_sample(sample: Dict, factor: int = 5) -> List[Dict]:
"""Generate augmented variations using templates + slang injection"""
augmented = [sample] # Keep original
sentiment = sample["sentiment"].lower()
asset = random.choice(ASSETS)
price = random.randint(100, 100000)
pct = random.randint(5, 50)
p_low = price - random.randint(10, 500)
p_high = price + random.randint(10, 500)
templates = CRYPTO_TEMPLATES.get(sentiment, CRYPTO_TEMPLATES["neutral"])
for _ in range(factor - 1):
template = random.choice(templates)
new_text = template.format(
asset=asset, price=price, pct=pct, p_low=p_low, p_high=p_high
)
# Inject random slang
if random.random() < 0.3:
slang = random.choice(CRYPTO_SLANG)
new_text = f"{new_text} {slang}!"
augmented.append({
"text": new_text,
"sentiment": sample["sentiment"],
"event_type": sample.get("event_type", "unknown"),
"entities": [asset],
})
return augmented
def build_dataset() -> List[Dict]:
"""Build training dataset from labeled files + augmentation"""
print("📦 Loading labeled data...")
all_data = []
for fname in ["labeled_verified.jsonl", "labeled_expanded.jsonl", "labeled_large.jsonl"]:
path = f"data/{fname}"
if os.path.exists(path):
loaded = load_labeled_data(path)
print(f" {fname}: {len(loaded)} samples")
all_data.extend(loaded)
# Deduplicate by text hash
seen = set()
unique = []
for d in all_data:
h = hashlib.md5(d["text"].encode()).hexdigest()[:16]
if h not in seen:
seen.add(h)
unique.append(d)
print(f" Unique samples: {len(unique)}")
# Augment
print(f"🔄 Augmenting x{CFG.aug_factor}...")
augmented = []
for sample in unique:
augmented.extend(augment_sample(sample, CFG.aug_factor))
print(f" Total after augmentation: {len(augmented)}")
return augmented
# ============================================================
# HUMAN-IN-THE-LOOP VERIFICATION
# ============================================================
@dataclass
class VerificationRecord:
sample_id: str
text: str
predicted: str
confidence: float
human_label: Optional[str] = None
verified_at: Optional[str] = None
verified_by: str = "cli"
class HumanVerifier:
"""CLI-based human verification with persistent state"""
def __init__(self, verify_file: str = "data/verification_log.jsonl"):
self.verify_file = Path(verify_file)
self.records: List[VerificationRecord] = []
self._load_existing()
def _load_existing(self):
if self.verify_file.exists():
for line in open(self.verify_file):
try:
self.records.append(VerificationRecord(**json.loads(line)))
except:
pass
print(f"📋 Loaded {len(self.records)} existing verifications")
def _save_record(self, record: VerificationRecord):
with open(self.verify_file, "a") as f:
f.write(json.dumps(asdict(record)) + "\n")
def verify_batch(self, samples: List[Dict], predictions: List[Tuple[str, float]]) -> List[Dict]:
"""Interactive verification of a batch"""
print(f"\n{'='*60}")
print(f"👁 HUMAN VERIFICATION — Batch of {len(samples)}")
print(f"{'='*60}")
verified = []
for i, (sample, (pred, conf)) in enumerate(zip(samples, predictions)):
print(f"\n--- Sample {i+1}/{len(samples)} (conf: {conf:.2%}) ---")
print(f"Text: {sample['text'][:200]}...")
print(f"Model: {pred} ({conf:.2%})")
# Quick keys: b=bullish, r=bearish, n=neutral, s=skip, q=quit
while True:
choice = input("Label [B]ullish/[R]earish/[N]eutral/[S]kip/[Q]uit: ").strip().lower()
if choice in ('b', 'bullish'):
human = "Bullish"
break
elif choice in ('r', 'bearish'):
human = "Bearish"
break
elif choice in ('n', 'neutral'):
human = "Neutral"
break
elif choice in ('s', 'skip'):
human = None
break
elif choice in ('q', 'quit'):
print("🛑 Quitting verification")
return verified
else:
print(" Invalid. Use B/R/N/S/Q")
if human:
record = VerificationRecord(
sample_id=hashlib.md5(sample["text"].encode()).hexdigest()[:12],
text=sample["text"],
predicted=pred,
confidence=conf,
human_label=human,
verified_at=datetime.now().isoformat(),
)
self.records.append(record)
self._save_record(record)
verified.append({**sample, "sentiment": human})
print(f" ✅ Saved: {human}")
else:
print(f" ⏭ Skipped")
return verified
# ============================================================
# LORA TRAINING
# ============================================================
def train_lora(train_data: List[Dict], val_data: List[Dict]) -> str:
"""Train LoRA adapter, return path to best model"""
if not TORCH_AVAILABLE:
print("❌ torch/peft not available. Cannot train.")
return ""
print(f"\n🚀 Starting LoRA training on {len(train_data)} samples...")
# Label mapping
label2id = {"Bearish": 0, "Bullish": 1, "Neutral": 2}
id2label = {v: k for k, v in label2id.items()}
# Prepare datasets
def to_dataset(data):
texts = [d["text"] for d in data]
labels = [label2id.get(d["sentiment"], 2) for d in data]
return Dataset.from_dict({"text": texts, "label": labels})
train_ds = to_dataset(train_data)
val_ds = to_dataset(val_data)
# Tokenizer
tokenizer = AutoTokenizer.from_pretrained(CFG.base_model)
def tokenize(batch):
return tokenizer(batch["text"], truncation=True, max_length=CFG.max_length, padding="max_length")
train_ds = train_ds.map(tokenize, batched=True)
val_ds = val_ds.map(tokenize, batched=True)
train_ds.set_format("torch", columns=["input_ids", "attention_mask", "label"])
val_ds.set_format("torch", columns=["input_ids", "attention_mask", "label"])
# Model + LoRA
model = AutoModelForSequenceClassification.from_pretrained(
CFG.base_model, num_labels=3, id2label=id2label, label2id=label2id
)
lora_config = LoraConfig(
r=CFG.r, lora_alpha=CFG.lora_alpha, lora_dropout=CFG.lora_dropout,
target_modules=list(CFG.target_modules), bias="none", task_type=TaskType.SEQ_CLS
)
model = get_peft_model(model, lora_config)
model.print_trainable_parameters()
# Training args (disk-conscious)
output_dir = CFG.output_dir
os.makedirs(output_dir, exist_ok=True)
training_args = TrainingArguments(
output_dir=output_dir,
num_train_epochs=CFG.num_epochs,
per_device_train_batch_size=CFG.batch_size,
per_device_eval_batch_size=CFG.batch_size * 2,
gradient_accumulation_steps=CFG.grad_accum,
learning_rate=CFG.learning_rate,
warmup_ratio=CFG.warmup_ratio,
weight_decay=CFG.weight_decay,
max_grad_norm=1.0,
eval_strategy="steps",
eval_steps=CFG.save_steps,
save_strategy="steps",
save_steps=CFG.save_steps,
save_total_limit=CFG.save_total_limit,
load_best_model_at_end=True,
metric_for_best_model="f1_macro",
greater_is_better=True,
fp16=False, # CPU
dataloader_num_workers=0,
logging_steps=10,
remove_unused_columns=False,
report_to="none",
seed=42,
)
def compute_metrics(eval_pred):
from sklearn.metrics import f1_score, accuracy_score
logits, labels = eval_pred
preds = logits.argmax(-1)
return {
"f1_macro": f1_score(labels, preds, average="macro"),
"accuracy": accuracy_score(labels, preds),
}
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_ds,
eval_dataset=val_ds,
tokenizer=tokenizer,
compute_metrics=compute_metrics,
)
print("🏋️ Training...")
trainer.train()
# Save best
best_path = os.path.join(output_dir, "best")
trainer.save_model(best_path)
tokenizer.save_pretrained(best_path)
# Cleanup old checkpoints (keep only best)
for d in os.listdir(output_dir):
if d.startswith("checkpoint-"):
shutil.rmtree(os.path.join(output_dir, d), ignore_errors=True)
print(f"✅ Best model saved to: {best_path}")
return best_path
# ============================================================
# MAIN PIPELINE
# ============================================================
def run_pipeline(mode: str = "train"):
"""Main entry point"""
print("="*60)
print("LoRA CRYPTO SENTIMENT PIPELINE")
print("="*60)
print(f"Mode: {mode}")
print(f"Config: r={CFG.r}, alpha={CFG.lora_alpha}, epochs={CFG.num_epochs}")
print(f"Disk limit: {CFG.max_checkpoint_gb}GB checkpoints")
# Build dataset
all_data = build_dataset()
random.shuffle(all_data)
# Split train/val
split = int(0.9 * len(all_data))
train_data = all_data[:split]
val_data = all_data[split:]
print(f"Train: {len(train_data)} | Val: {len(val_data)}")
if mode == "train":
# Train
best_path = train_lora(train_data, val_data)
print(f"\n🎉 Training complete! Best model: {best_path}")
elif mode == "verify":
# Human verification mode - run model on new data, verify
print("🔮 Loading model for verification...")
# This would load the trained LoRA and run on new sources
print("Run: python lora_pipeline.py verify --model-path <path>")
elif mode == "augment-only":
# Just show augmented samples
for d in train_data[:5]:
print(f"\n{d['sentiment']}: {d['text'][:100]}...")
print("\n✅ Pipeline complete")
# ============================================================
# CLI
# ============================================================
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(description="LoRA Crypto Sentiment Pipeline")
parser.add_argument("mode", choices=["train", "verify", "augment-only"], default="train")
parser.add_argument("--epochs", type=int, default=3)
parser.add_argument("--r", type=int, default=8)
parser.add_argument("--batch", type=int, default=8)
parser.add_argument("--lr", type=float, default=2e-4)
parser.add_argument("--max-len", type=int, default=128)
parser.add_argument("--aug", type=int, default=5)
parser.add_argument("--output", type=str, default="./models/lora-finbert-crypto")
args = parser.parse_args()
# Override config
CFG.num_epochs = args.epochs
CFG.r = args.r
CFG.batch_size = args.batch
CFG.learning_rate = args.lr
CFG.max_length = args.max_len
CFG.aug_factor = args.aug
CFG.output_dir = args.output
run_pipeline(args.mode)