Add sentiment_engine with CryptoSentimentCalibrator fixes - improved keyword lists, lowered FinBERT threshold, added neutral handling
This commit is contained in:
300
sentiment_engine/training/improve_sentiment_model.py
Normal file
300
sentiment_engine/training/improve_sentiment_model.py
Normal file
@@ -0,0 +1,300 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Improve sentiment model with more training data and better training.
|
||||
Uses labeled_verified.jsonl + augmented data from specs.
|
||||
"""
|
||||
|
||||
import json
|
||||
import random
|
||||
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
|
||||
)
|
||||
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
|
||||
|
||||
SENTIMENT_LABELS = ["Bearish", "Bullish", "Neutral"]
|
||||
SENTIMENT_MAP = {"Bearish": 0, "Bullish": 1, "Neutral": 2}
|
||||
|
||||
def load_labeled_data(label_file: str):
|
||||
"""Load verified labeled data from labeling pipeline"""
|
||||
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():
|
||||
"""Build comprehensive training data from Spec #2 keywords + labeled data"""
|
||||
|
||||
# Bearish samples (from Spec #2 bearish keywords + verified data)
|
||||
bearish_texts = [
|
||||
# From labeled verified data
|
||||
"Major hack: Radiant Capital loses $50M in exploit. Attacker exploits rounding error in lending market. Funds moved to Tornado Cash.",
|
||||
"Curve Finance hit by $50M exploit. Vyper compiler bug affects multiple pools. CRV drops 20%.",
|
||||
"Wintermute market maker loses $20M in exploit. Private key compromise suspected. Funds returned.",
|
||||
"SEC sues Kraken for operating unregistered securities exchange. Alleged commingling of customer funds.",
|
||||
"SEC charges Uniswap Labs with operating unregistered securities exchange. UNI drops 15%.",
|
||||
"Binance delists Monero (XMR), Zcash (ZEC), and 4 other privacy coins. Cites regulatory compliance review.",
|
||||
"OKX delists USDT trading pairs in EEA region. MiCA compliance cited. USDT/USD pairs remain.",
|
||||
"Solana network experiences 5-hour outage. Validators restart cluster. SOL drops 8% on news.",
|
||||
"Circle USDC depegs to $0.97 after SVB exposure revealed. $3.3B reserves stuck at SVB. Arbitrage bots profit.",
|
||||
# Spec #2 bearish keywords expanded
|
||||
"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",
|
||||
]
|
||||
|
||||
# Bullish samples
|
||||
bullish_texts = [
|
||||
# From labeled verified data
|
||||
"Bitcoin hits new all-time high of $108,000 as institutional inflows surge. BlackRock IBIT ETF sees record $1.2B daily inflow.",
|
||||
"Bitcoin ETF inflows hit record $2.1B in single week. IBIT alone sees $1.2B. Cumulative AUM passes $50B.",
|
||||
"Bitcoin hits $100,000 for first time ever. MicroStrategy, ETFs, and 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. Stock MSTR up 15% premarket.",
|
||||
"Arbitrum DAO approves $200M ARB grant program for gaming ecosystem. Voting passes with 92% approval.",
|
||||
"EigenLayer restaking TVL hits $20B. ETH restaking becomes largest DeFi category. Points season 2 announced.",
|
||||
"Hyperliquid DEX launches HYPE token airdrop. $1.2B TVL locked. Points program drives volume.",
|
||||
"dYdX chain migration to Cosmos complete. V4 mainnet launches with 0.02s block times. DYDX token migration.",
|
||||
# Spec #2 bullish keywords expanded
|
||||
"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",
|
||||
]
|
||||
|
||||
# Neutral samples
|
||||
neutral_texts = [
|
||||
# From labeled verified data
|
||||
"Ethereum Dencun upgrade goes live on mainnet. Proto-Danksharding (EIP-4844) activates, reducing L2 transaction fees by 90%.",
|
||||
"SEC approves spot Bitcoin ETFs for 11 issuers including BlackRock, Fidelity, ARK. Trading begins Thursday.",
|
||||
"Coinbase lists Pepe (PEPE) and Bonk (BONK) memecoins. Trading opens with 100x volume spike.",
|
||||
"Ethereum Pectra upgrade activated. EIP-7702 account abstraction live. EOAs can now batch transactions.",
|
||||
# Spec #2 neutral/descriptive keywords
|
||||
"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",
|
||||
]
|
||||
|
||||
# Build training data
|
||||
texts = []
|
||||
labels = []
|
||||
|
||||
for t in bearish_texts:
|
||||
texts.append(t); labels.append(0)
|
||||
for t in bullish_texts:
|
||||
texts.append(t); labels.append(1)
|
||||
for t in neutral_texts:
|
||||
texts.append(t); labels.append(2)
|
||||
|
||||
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("IMPROVING SENTIMENT MODEL - EXPANDED TRAINING")
|
||||
print("="*60)
|
||||
|
||||
# 1. Load verified labeled data
|
||||
print("\n1. Loading verified labeled data...")
|
||||
verified_texts, verified_labels = load_labeled_data("data/labeled_verified.jsonl")
|
||||
print(f" Verified samples: {len(verified_texts)}")
|
||||
|
||||
# 2. Build augmented data from Spec #2
|
||||
print("\n2. Building augmented training data from Spec #2...")
|
||||
aug_texts, aug_labels = build_augmented_data()
|
||||
print(f" Augmented samples: {len(aug_texts)}")
|
||||
|
||||
# 3. Combine (weight verified data higher by duplicating)
|
||||
all_texts = verified_texts * 3 + aug_texts # 3x weight for verified
|
||||
all_labels = verified_labels * 3 + aug_labels
|
||||
|
||||
print(f"\n3. Total training samples: {len(all_texts)}")
|
||||
print(f" Bearish: {all_labels.count(0)}, Bullish: {all_labels.count(1)}, Neutral: {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 tokenizer and model
|
||||
print("\n4. Loading model...")
|
||||
tokenizer = AutoTokenizer.from_pretrained("models/finbert-crypto-sentiment")
|
||||
model = AutoModelForSequenceClassification.from_pretrained("models/finbert-crypto-sentiment")
|
||||
|
||||
# 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 for balanced training
|
||||
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
|
||||
training_args = TrainingArguments(
|
||||
output_dir="./models/finbert-crypto-sentiment-v2",
|
||||
num_train_epochs=4,
|
||||
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=10,
|
||||
save_total_limit=2,
|
||||
remove_unused_columns=False,
|
||||
report_to="none",
|
||||
weight_decay=0.01,
|
||||
)
|
||||
|
||||
# Custom trainer with class weights
|
||||
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=2)]
|
||||
)
|
||||
|
||||
print("\n5. Training (4 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", 1), # Bullish
|
||||
("Bitcoin crashes 50% in hours, massive selloff", 0), # Bearish
|
||||
("BTC at $50k, ETH at $3k, market consolidating", 2), # Neutral
|
||||
]
|
||||
|
||||
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("SENTIMENT MODEL IMPROVEMENT COMPLETE!")
|
||||
print("="*60)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user