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

315 lines
14 KiB
Python
Raw Normal View History

#!/usr/bin/env python3
"""
Retrain sentiment model for CRYPTO - flip Bearish/Bullish to match crypto semantics.
FinBERT native: 0=negative, 1=neutral, 2=positive
Crypto mapping: "surge/moon/pump" -> Bullish(2), "crash/dump/rug" -> Bearish(0)
"""
import json
import torch
import numpy as np
from pathlib import Path
from typing import List
from torch.utils.data import Dataset
from transformers import (
AutoTokenizer, AutoModelForSequenceClassification,
TrainingArguments, Trainer, EarlyStoppingCallback
)
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score, f1_score
from sklearn.utils.class_weight import compute_class_weight
# CRYPTO label mapping - matches how crypto traders think
# 0 = Bearish (price down, crash, dump, hack, rug)
# 1 = Neutral (consolidation, upgrade, listing, regulatory)
# 2 = Bullish (price up, surge, pump, moon, inflow, adoption)
SENTIMENT_LABELS = ["Bearish", "Neutral", "Bullish"]
SENTIMENT_MAP = {"Bearish": 0, "Neutral": 1, "Bullish": 2}
def load_labeled_data(label_file: str):
"""Load verified labeled data - FLIP bearish/bullish for crypto"""
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'])
orig_label = r['labels']['sentiment']
# FinBERT labeled these with traditional finance labels
# For crypto, we need to FLIP: Bearish<->Bullish
if orig_label == "Bearish":
labels.append(2) # Flip to Bullish (FinBERT's positive)
elif orig_label == "Bullish":
labels.append(0) # Flip to Bearish (FinBERT's negative)
else:
labels.append(1) # Neutral stays Neutral
return texts, labels
def build_augmented_data():
"""Build training data with CRYPTO semantics"""
# 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.",
]
# 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",
]
# 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%.",
"SEC approves spot Bitcoin ETFs for 11 issuers. Trading begins Thursday.",
"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 chain migration to Cosmos complete. V4 mainnet launches.",
]
texts = []
labels = []
for t in bearish_texts:
texts.append(t); labels.append(0) # Bearish = 0
for t in bullish_texts:
texts.append(t); labels.append(2) # Bullish = 2
for t in neutral_texts:
texts.append(t); labels.append(1) # Neutral = 1
return texts, labels
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()
}
def main():
print("="*60)
print("RETRAINING SENTIMENT FOR CRYPTO SEMANTICS")
print("="*60)
# 1. Load verified labeled data (FLIPPED)
print("\n1. Loading verified labeled data (with label flip)...")
verified_texts, verified_labels = load_labeled_data("data/labeled_verified.jsonl")
print(f" Verified samples: {len(verified_texts)}")
# 2. Build augmented data with crypto semantics
print("\n2. Building augmented training data (crypto semantics)...")
aug_texts, aug_labels = build_augmented_data()
print(f" Augmented samples: {len(aug_texts)}")
# 3. Combine (weight verified 3x)
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
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)}")
# 5. Load BASE FinBERT
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 - heavily weight Bearish since it's hardest
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)
# Boost Bearish weight further
class_weights[0] *= 2.0
print(f" Class weights: {class_weights}")
# 8. Training arguments
training_args = TrainingArguments(
output_dir="./models/finbert-crypto-sentiment-v3",
num_train_epochs=6,
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=2,
remove_unused_columns=False,
report_to="none",
weight_decay=0.01,
)
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)]
)
print("\n5. Training (6 epochs)...")
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. Save best model
print("\n7. Saving improved model...")
model.save_pretrained("./models/finbert-crypto-sentiment")
tokenizer.save_pretrained("./models/finbert-crypto-sentiment")
print(" ✅ Model saved to models/finbert-crypto-sentiment/")
# 11. Quick inference test
print("\n8. Quick inference test...")
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),
]
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]
print(f" '{text[:50]}...' -> {SENTIMENT_LABELS[pred]} (polarity={polarity:.2f}) expected={SENTIMENT_LABELS[expected]}")
print("\n" + "="*60)
print("CRYPTO SENTIMENT MODEL TRAINING COMPLETE!")
print("="*60)
if __name__ == "__main__":
main()