diff --git a/sentiment_engine/training/lora_pipeline.py b/sentiment_engine/training/lora_pipeline.py new file mode 100644 index 0000000..d32c4ca --- /dev/null +++ b/sentiment_engine/training/lora_pipeline.py @@ -0,0 +1,475 @@ +#!/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 ") + + 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) diff --git a/sentiment_engine/training/requirements-lora.txt b/sentiment_engine/training/requirements-lora.txt new file mode 100644 index 0000000..923c99c --- /dev/null +++ b/sentiment_engine/training/requirements-lora.txt @@ -0,0 +1,8 @@ +# Minimal requirements for LoRA pipeline (CPU-friendly) +torch==2.3.0 --index-url https://download.pytorch.org/whl/cpu +transformers==4.44.0 +peft==0.12.0 +datasets==2.20.0 +accelerate==0.33.0 +scikit-learn==1.5.0 +numpy==1.26.0