357 lines
16 KiB
Python
357 lines
16 KiB
Python
#!/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()
|