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

357 lines
16 KiB
Python
Raw Normal View History

#!/usr/bin/env python3
"""
Robust FinBERT fine-tuning for crypto sentiment with:
- Expanded labeled data (92 verified samples)
- Proper crypto semantics (Bearish=0, Neutral=1, Bullish=2)
- Checkpoint-based training to prevent forgetting
- Class-weighted loss, early stopping, LR scheduling
- Saves best model based on validation F1_macro
"""
import json
import torch
import numpy as np
from pathlib import Path
from typing import List, Dict
from torch.utils.data import Dataset
from transformers import (
AutoTokenizer, AutoModelForSequenceClassification,
TrainingArguments, Trainer, EarlyStoppingCallback, TrainerCallback
)
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score, f1_score, classification_report
from sklearn.utils.class_weight import compute_class_weight
import logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
# CRYPTO SENTIMENT LABELS (matches FinBERT native order: negative=0, neutral=1, positive=2)
# For crypto: Bearish(negative)=0, Neutral=1, Bullish(positive)=2
SENTIMENT_LABELS = ["Bearish", "Neutral", "Bullish"]
SENTIMENT_MAP = {"Bearish": 0, "Neutral": 1, "Bullish": 2}
class SentimentDataset(Dataset):
def __init__(self, texts, labels, tokenizer, max_len=128):
self.texts = texts
self.labels = labels
self.tokenizer = tokenizer
self.max_len = max_len
def __len__(self):
return len(self.texts)
def __getitem__(self, idx):
text = self.texts[idx]
label = self.labels[idx]
encoding = self.tokenizer(
text,
truncation=True,
max_length=self.max_len,
padding="max_length",
return_tensors="pt"
)
return {
"input_ids": encoding["input_ids"].squeeze(0),
"attention_mask": encoding["attention_mask"].squeeze(0),
"token_type_ids": encoding.get("token_type_ids", torch.zeros_like(encoding["input_ids"])).squeeze(0),
"labels": torch.tensor(label, dtype=torch.long)
}
def compute_metrics(eval_pred):
predictions, labels = eval_pred
predictions = np.argmax(predictions, axis=1)
return {
"accuracy": accuracy_score(labels, predictions),
"f1_macro": f1_score(labels, predictions, average="macro"),
"f1_per_class": f1_score(labels, predictions, average=None).tolist()
}
class BestModelCheckpoint(TrainerCallback):
"""Custom callback to save best model based on validation F1_macro"""
def __init__(self, save_path: str):
self.save_path = save_path
self.best_f1 = 0.0
def on_evaluate(self, args, state, control, metrics=None, **kwargs):
if metrics is not None:
eval_f1 = metrics.get("eval_f1_macro", 0)
if eval_f1 > self.best_f1:
self.best_f1 = eval_f1
logger.info(f"New best F1_macro: {eval_f1:.4f} - saving model to {self.save_path}")
# The Trainer handles saving via save_strategy="epoch" and load_best_model_at_end=True
# This callback just tracks the best metric
def load_labeled_data(label_file: str):
"""Load verified labeled data from labeling pipeline output"""
texts, labels = [], []
with open(label_file) as f:
for line in f:
r = json.loads(line)
if r.get('verified', False):
texts.append(r['text'])
labels.append(SENTIMENT_MAP[r['labels']['sentiment']])
return texts, labels
def build_augmented_data():
"""Additional high-quality synthetic samples for data augmentation"""
# CRYPTO BEARISH (label 0) - price down, bad news
bearish_texts = [
"Bitcoin crashes 30% in hours as leverage flushes out longs",
"Massive liquidation cascade wipes out $500M in longs across exchanges",
"Exchange hacked, $100M stolen, users panic selling",
"Regulatory crackdown: SEC files enforcement action against major DeFi protocol",
"Rug pull: Dev team abandons project, drains liquidity pool",
"Bankruptcy filing: Major crypto lender files Chapter 11",
"Stablecoin depeg: USDT drops to $0.95 on redemption fears",
"Smart contract vulnerability discovered, $50M at risk",
"Market structure breakdown: Order books thin, spreads widen",
"Forced liquidations trigger death spiral in lending protocol",
"Contagion risk: Major fund exposure to failed protocol revealed",
"Bear market confirmed: Lower highs, lower lows on weekly chart",
"Institutional outflows: ETF sees record redemptions for 5th week",
"Mining capitulation: Hash rate drops 20% as price falls below cost",
"Major hack: Radiant Capital loses $50M in exploit. Funds moved to Tornado Cash.",
"Curve Finance hit by $50M exploit. Vyper compiler bug. CRV drops 20%.",
"Wintermute market maker loses $20M in exploit. Funds returned.",
"SEC sues Kraken for operating unregistered securities exchange.",
"SEC charges Uniswap Labs. UNI drops 15%.",
"Binance delists Monero, Zcash, and 4 other privacy coins.",
"OKX delists USDT trading pairs in EEA region. MiCA compliance.",
"Solana network experiences 5-hour outage. SOL drops 8% on news.",
"Circle USDC depegs to $0.97 after SVB exposure. $3.3B reserves stuck.",
"Australia ASIC cracks down on unlicensed crypto exchanges.",
"Treasury Secretary Yellen comments on stablecoin regulation.",
]
# CRYPTO BULLISH (label 2) - price up, good news
bullish_texts = [
"Bitcoin surges to $108k as institutional inflows surge. BlackRock IBIT sees record $1.2B daily inflow.",
"Bitcoin ETF inflows hit record $2.1B in single week. Cumulative AUM passes $50B.",
"Bitcoin hits $100,000 for first time ever. MicroStrategy, ETFs, sovereign buying drive rally.",
"Pump.fun revenue hits $100M in 30 days. Memecoin factory launches 50k tokens/day. SOL fees surge.",
"MicroStrategy buys additional 12,000 BTC at $61M. Total holdings now 190,000 BTC.",
"Institutional adoption accelerates: Fortune 500 companies adding BTC to treasury",
"ETF approval drives massive inflows: $10B in first month",
"Golden cross confirmed on Bitcoin weekly chart, technical breakout",
"Supply shock: Exchange balances hit 5-year low as holders accumulate",
"Layer 2 adoption surges: Arbitrum and Optimism TVL doubles",
"Real yield protocols attract TradFi capital seeking returns",
"Token unlock schedule favorable: Low float, high demand dynamics",
"Major partnership: TradFi giant integrates blockchain settlement",
"Sovereign wealth fund announces Bitcoin allocation",
"Hash rate hits all-time high, mining investment surges",
"Developer activity reaches record highs across major ecosystems",
"Stablecoin supply grows 50% YoY, indicating fresh capital entry",
"Options market signaling upside: Call skew at multi-year highs",
"Macro tailwinds: Rate cuts expected, dollar weakening",
"SEC approves spot Bitcoin ETFs for 11 issuers including BlackRock, Fidelity, ARK.",
"Hong Kong SFC approves spot Bitcoin ETFs. Asia ETF race begins.",
"Canada OSC approves first spot Solana ETF. North American product expansion.",
"Safe{Wallet} hits $100B secured. Multi-sig adoption standard.",
"Ripple wins landmark court case against SEC. XRP surges 70%.",
]
# CRYPTO NEUTRAL (label 1) - sideways, structural, non-directional
neutral_texts = [
"Bitcoin consolidates in tight range between $50k-$52k",
"Ethereum gas fees stable at 15 gwei amid low activity",
"Market awaits FOMC decision, volumes below average",
"Trading range established: Support at $48k, resistance at $55k",
"Altcoin season index neutral at 50, no clear trend",
"Funding rates flat across perpetual futures markets",
"On-chain metrics show equilibrium: Inflows match outflows",
"Derivatives open interest stable, no excessive leverage",
"Stablecoin market cap flat month-over-month",
"Developer conference announces roadmap, no token news",
"Governance proposal passes: Parameter change only, no value accrual",
"Exchange lists new token, volume modest, no price impact",
"Research report: Fair value estimate $55k-$65k range",
"Whale wallet rotates positions, no net accumulation or distribution",
"Ethereum Dencun upgrade goes live. Proto-Danksharding reduces L2 fees by 90%.",
"Coinbase lists Pepe and Bonk memecoins. Trading opens with 100x volume spike.",
"Ethereum Pectra upgrade activated. EIP-7702 account abstraction live.",
"Arbitrum DAO approves $200M ARB grant program. Voting passes with 92%.",
"EigenLayer restaking TVL hits $20B. Points season 2 announced.",
"Hyperliquid DEX launches HYPE token airdrop. $1.2B TVL locked.",
"dYdX V4 mainnet launches. Cosmos-based order book DEX.",
"EigenLayer restaking TVL hits $25B. Largest DeFi category.",
"Ripple wins landmark court case against SEC. XRP surges 70% on ruling.",
"Babylon Bitcoin staking testnet. Bitcoin security for PoS chains.",
"Safe{Wallet} hits $100B secured. Multi-sig adoption standard.",
]
texts = []
labels = []
for t in bearish_texts:
texts.append(t); labels.append(0)
for t in bullish_texts:
texts.append(t); labels.append(2)
for t in neutral_texts:
texts.append(t); labels.append(1)
return texts, labels
def main():
print("="*60)
print("ROBUST FINBERT FINE-TUNING FOR CRYPTO SENTIMENT")
print("="*60)
# 1. Load verified labeled data (primary source)
print("\n1. Loading verified labeled data from labeling pipeline...")
verified_texts, verified_labels = load_labeled_data("data/labeled_verified.jsonl")
print(f" Verified samples: {len(verified_texts)}")
# 2. Build augmented data
print("\n2. Building augmented training data...")
aug_texts, aug_labels = build_augmented_data()
print(f" Augmented samples: {len(aug_texts)}")
# 3. Combine with weighted emphasis on verified data (3x weight)
all_texts = verified_texts * 3 + aug_texts
all_labels = verified_labels * 3 + aug_labels
print(f"\n3. Total training samples: {len(all_texts)}")
print(f" Bearish(0): {all_labels.count(0)}, Neutral(1): {all_labels.count(1)}, Bullish(2): {all_labels.count(2)}")
# 4. Train/val/test split (stratified)
train_texts, temp_texts, train_labels, temp_labels = train_test_split(
all_texts, all_labels, test_size=0.3, random_state=42, stratify=all_labels
)
val_texts, test_texts, val_labels, test_labels = train_test_split(
temp_texts, temp_labels, test_size=0.5, random_state=42, stratify=temp_labels
)
print(f" Train: {len(train_texts)}, Val: {len(val_texts)}, Test: {len(test_texts)}")
print(f" Train dist: Bearish={train_labels.count(0)}, Neutral={train_labels.count(1)}, Bullish={train_labels.count(2)}")
# 5. Load BASE FinBERT (not previously fine-tuned)
print("\n4. Loading BASE FinBERT...")
tokenizer = AutoTokenizer.from_pretrained("ProsusAI/finbert")
model = AutoModelForSequenceClassification.from_pretrained(
"ProsusAI/finbert",
num_labels=3,
id2label={0: "Bearish", 1: "Neutral", 2: "Bullish"},
label2id={"Bearish": 0, "Neutral": 1, "Bullish": 2}
)
# 6. Create datasets
train_dataset = SentimentDataset(train_texts, train_labels, tokenizer, max_len=128)
val_dataset = SentimentDataset(val_texts, val_labels, tokenizer, max_len=128)
# 7. Class weights - compute from training data
class_weights = compute_class_weight("balanced", classes=np.array([0,1,2]), y=np.array(train_labels))
class_weights = torch.tensor(class_weights, dtype=torch.float)
print(f" Class weights: {class_weights}")
# 8. Training arguments with robust settings
output_dir = "./models/finbert-crypto-sentiment-v4"
training_args = TrainingArguments(
output_dir=output_dir,
num_train_epochs=8,
per_device_train_batch_size=8,
per_device_eval_batch_size=16,
gradient_accumulation_steps=4,
warmup_ratio=0.1,
learning_rate=1e-5,
lr_scheduler_type="cosine",
eval_strategy="epoch",
save_strategy="epoch",
load_best_model_at_end=True,
metric_for_best_model="f1_macro",
greater_is_better=True,
fp16=False,
dataloader_num_workers=0,
logging_steps=5,
save_total_limit=3,
remove_unused_columns=False,
report_to="none",
weight_decay=0.01,
max_grad_norm=1.0,
)
class WeightedTrainer(Trainer):
def compute_loss(self, model, inputs, return_outputs=False, **kwargs):
labels = inputs.get("labels")
outputs = model(**inputs)
logits = outputs.get("logits")
loss_fct = torch.nn.CrossEntropyLoss(weight=class_weights.to(logits.device))
loss = loss_fct(logits.view(-1, 3), labels.view(-1))
return (loss, outputs) if return_outputs else loss
trainer = WeightedTrainer(
model=model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=val_dataset,
tokenizer=tokenizer,
compute_metrics=compute_metrics,
callbacks=[
EarlyStoppingCallback(early_stopping_patience=3),
BestModelCheckpoint(output_dir)
]
)
print("\n5. Training (8 epochs with early stopping)...")
trainer.train()
# 9. Evaluate on test set
print("\n6. Evaluating on test set...")
test_dataset = SentimentDataset(test_texts, test_labels, tokenizer, max_len=128)
test_results = trainer.evaluate(test_dataset)
print(f" Test results: {test_results}")
# 10. Detailed classification report
print("\n7. Detailed classification report...")
test_trainer = Trainer(model=model, tokenizer=tokenizer, compute_metrics=compute_metrics)
predictions = test_trainer.predict(test_dataset)
preds = np.argmax(predictions.predictions, axis=1)
print(classification_report(test_labels, preds, target_names=SENTIMENT_LABELS))
# 11. Save best model to production path
print("\n8. Saving production model...")
model.save_pretrained("./models/finbert-crypto-sentiment")
tokenizer.save_pretrained("./models/finbert-crypto-sentiment")
print(" ✅ Model saved to models/finbert-crypto-sentiment/")
# 12. Quick inference test on critical cases
print("\n9. Quick inference test on critical cases...")
model.eval()
test_cases = [
("Bitcoin surges to $108k as institutional inflows surge", 2),
("Bitcoin crashes 50% in hours, massive selloff", 0),
("BTC at $50k, ETH at $3k, market consolidating", 1),
("Major hack on exchange, $100M stolen, panic selling", 0),
("ETF approval drives massive inflows, price to moon", 2),
("Market consolidating in tight range, no clear direction", 1),
("Circle USDC depegs to $0.97 after SVB exposure", 0),
("SEC sues Kraken for operating unregistered securities", 0),
("SEC approves spot Bitcoin ETFs for 11 issuers", 2),
("Australia ASIC cracks down on unlicensed exchanges", 0),
]
print(f"{'Text':<60} {'Pred':<10} {'Exp':<10} {'Polarity':<10} {'Conf':<6}")
print("-" * 100)
for text, expected in test_cases:
inputs = tokenizer(text, return_tensors="pt", truncation=True, max_length=128, padding=True)
with torch.no_grad():
outputs = model(**inputs)
probs = torch.softmax(outputs.logits, dim=-1).numpy()[0]
pred = np.argmax(probs)
polarity = probs[2] - probs[0]
conf = probs[pred]
status = "✓" if pred == expected else "✗"
print(f"{text[:58]:<60} {SENTIMENT_LABELS[pred]:<10} {SENTIMENT_LABELS[expected]:<10} {polarity:>+6.2f} {conf:.2f} {status}")
print("\n" + "="*60)
print("FINE-TUNING COMPLETE!")
print("="*60)
if __name__ == "__main__":
main()