397 lines
16 KiB
Python
397 lines
16 KiB
Python
|
|
#!/usr/bin/env python3
|
||
|
|
"""
|
||
|
|
LoRA Fine-Tuning Pipeline for Crypto Emotion (Multi-Label)
|
||
|
|
- Base: j-hartmann/emotion-english-distilroberta-base (6 emotions)
|
||
|
|
- Multi-label classification (greed, fear, joy, anger, sadness, neutral)
|
||
|
|
- Crypto-specific emotion weighting (greed/fear emphasized)
|
||
|
|
- LoRA (r=8, alpha=16) on DistilRoBERTa
|
||
|
|
- Human-in-the-loop verification
|
||
|
|
"""
|
||
|
|
|
||
|
|
import os
|
||
|
|
import json
|
||
|
|
import random
|
||
|
|
import shutil
|
||
|
|
from pathlib import Path
|
||
|
|
from typing import List, Dict, Any
|
||
|
|
from dataclasses import dataclass
|
||
|
|
|
||
|
|
try:
|
||
|
|
import torch
|
||
|
|
import numpy as np
|
||
|
|
from transformers import AutoTokenizer, AutoModelForSequenceClassification, TrainingArguments, Trainer, EarlyStoppingCallback
|
||
|
|
from peft import LoraConfig, get_peft_model, TaskType
|
||
|
|
from datasets import Dataset
|
||
|
|
TORCH_AVAILABLE = True
|
||
|
|
except ImportError:
|
||
|
|
TORCH_AVAILABLE = False
|
||
|
|
print("⚠️ torch/peft not available")
|
||
|
|
|
||
|
|
# ============================================================
|
||
|
|
# CONFIG
|
||
|
|
# ============================================================
|
||
|
|
@dataclass
|
||
|
|
class EmotionLoRAConfig:
|
||
|
|
base_model: str = "j-hartmann/emotion-english-distilroberta-base"
|
||
|
|
output_dir: str = "./models/lora-distilroberta-crypto-emotion"
|
||
|
|
|
||
|
|
r: int = 8
|
||
|
|
lora_alpha: int = 16
|
||
|
|
lora_dropout: float = 0.1
|
||
|
|
target_modules: tuple = ("query", "key", "value", "intermediate.dense", "output.dense")
|
||
|
|
|
||
|
|
num_epochs: int = 5
|
||
|
|
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
|
||
|
|
save_steps: int = 50
|
||
|
|
save_total_limit: int = 1
|
||
|
|
|
||
|
|
emotion_labels: list = None
|
||
|
|
crypto_weights: dict = None
|
||
|
|
|
||
|
|
def __post_init__(self):
|
||
|
|
if self.emotion_labels is None:
|
||
|
|
self.emotion_labels = ["joy", "fear", "anger", "greed", "sadness", "neutral"]
|
||
|
|
if self.crypto_weights is None:
|
||
|
|
self.crypto_weights = {l: 1.0 for l in self.emotion_labels}
|
||
|
|
self.crypto_weights["greed"] = 2.0
|
||
|
|
self.crypto_weights["fear"] = 2.0
|
||
|
|
self.crypto_weights["joy"] = 1.5
|
||
|
|
|
||
|
|
CFG_EMO = EmotionLoRAConfig()
|
||
|
|
|
||
|
|
# ============================================================
|
||
|
|
# EMOTION AUGMENTATION TEMPLATES
|
||
|
|
# ============================================================
|
||
|
|
EMOTION_TEMPLATES = {
|
||
|
|
"greed": [
|
||
|
|
"FOMO driving {asset} to ${price}, apes accumulating",
|
||
|
|
"Buy the dip on {asset}! Loading bags at ${price}",
|
||
|
|
"Whale buying {asset} aggressively at ${price}",
|
||
|
|
"All in on {asset}! Diamond hands to ${price}",
|
||
|
|
"{asset} mooning! Greed index extreme, still buying",
|
||
|
|
"Accumulating {asset} heavily at ${price}, easy gains",
|
||
|
|
"Leverage long {asset} at ${price}, targeting ${target}",
|
||
|
|
"Smart money rotating into {asset} at ${price}",
|
||
|
|
"FOMO kicking in for {asset} at ${price}, don't miss out",
|
||
|
|
"{asset} breakout confirmed, greed taking over at ${price}",
|
||
|
|
],
|
||
|
|
"fear": [
|
||
|
|
"Panic selling {asset} at ${price}, liquidation cascade",
|
||
|
|
"Major hack drains {asset} liquidity, price crashes to ${price}",
|
||
|
|
"SEC crackdown sends {asset} plummeting {pct}%",
|
||
|
|
"Whale dumping {asset}, massive sell wall at ${price}",
|
||
|
|
"Support broken on {asset} at ${price}, fear index extreme",
|
||
|
|
"Rug pull suspected! {asset} collapsing to ${price}",
|
||
|
|
"Margin calls wiping {asset} longs at ${price}",
|
||
|
|
"Regulatory FUD crushing {asset} down to ${price}",
|
||
|
|
"Exchange halts {asset} withdrawals, panic at ${price}",
|
||
|
|
"Black swan event hits {asset}, price freefalls to ${price}",
|
||
|
|
],
|
||
|
|
"joy": [
|
||
|
|
"{asset} hits new ATH at ${price}! To the moon!",
|
||
|
|
"ETF approved for {asset}! Celebration at ${price}",
|
||
|
|
"Massive gains on {asset}! Diamond hands vindicated at ${price}",
|
||
|
|
"Breakthrough upgrade for {asset}! Joy at ${price}",
|
||
|
|
"Community celebrating {asset} milestone at ${price}",
|
||
|
|
"Whale buys {asset}! Euphoria at ${price}",
|
||
|
|
"Partnership announced for {asset}! Joy at ${price}",
|
||
|
|
"Mainnet launch successful for {asset} at ${price}",
|
||
|
|
"Record volume on {asset}! Party at ${price}",
|
||
|
|
"All-time high for {asset}! Pure joy at ${price}",
|
||
|
|
],
|
||
|
|
"anger": [
|
||
|
|
"Rug pull on {asset}! Devs drained liquidity!",
|
||
|
|
"Exchange froze {asset} withdrawals again! Furious!",
|
||
|
|
"Market manipulation on {asset}! Scam at ${price}",
|
||
|
|
"Insider trading suspected on {asset}! Anger at ${price}",
|
||
|
|
"Failed upgrade on {asset}! Incompetence at ${price}",
|
||
|
|
"Liquidity pulled on {asset}! Rage at ${price}",
|
||
|
|
"False promises from {asset} team! Betrayal at ${price}",
|
||
|
|
"Coordinated dump on {asset}! Market abuse at ${price}",
|
||
|
|
"Smart contract exploit on {asset}! Outrage at ${price}",
|
||
|
|
"Censorship on {asset} network! Freedom violation at ${price}",
|
||
|
|
],
|
||
|
|
"sadness": [
|
||
|
|
"Lost life savings on {asset} crash to ${price}",
|
||
|
|
"Bag holder on {asset}... down 90% from ATH",
|
||
|
|
"Missed the {asset} bottom... regret at ${price}",
|
||
|
|
"Rekt on {asset} leverage... pain at ${price}",
|
||
|
|
"Community devastated by {asset} rug pull",
|
||
|
|
"Years of gains wiped out on {asset} at ${price}",
|
||
|
|
"Trusted the {asset} team... betrayed at ${price}",
|
||
|
|
"Paper hands sold {asset} bottom... regret at ${price}",
|
||
|
|
"Dream of financial freedom crushed by {asset}",
|
||
|
|
"Watching {asset} bleed daily... hopeless at ${price}",
|
||
|
|
],
|
||
|
|
"neutral": [
|
||
|
|
"{asset} consolidating at ${price}, no clear direction",
|
||
|
|
"Low volume on {asset} at ${price}, waiting for catalyst",
|
||
|
|
"Market choppy for {asset} at ${price}",
|
||
|
|
"Sideways action on {asset} at ${price}",
|
||
|
|
"{asset} at ${price} with mixed signals",
|
||
|
|
"Accumulation phase for {asset} around ${price}",
|
||
|
|
"No news on {asset}, price stable at ${price}",
|
||
|
|
"Range-bound trading for {asset} at ${price}",
|
||
|
|
"Low volatility on {asset} at ${price}",
|
||
|
|
"Wait and see mode for {asset} at ${price}",
|
||
|
|
],
|
||
|
|
}
|
||
|
|
|
||
|
|
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", "ICP", "QNT", "INJ"]
|
||
|
|
|
||
|
|
# ============================================================
|
||
|
|
# BUILD EMOTION DATASET
|
||
|
|
# ============================================================
|
||
|
|
def build_emotion_dataset() -> List[Dict]:
|
||
|
|
print("🎭 Building emotion dataset...")
|
||
|
|
data = []
|
||
|
|
|
||
|
|
for emotion, templates in EMOTION_TEMPLATES.items():
|
||
|
|
for template in templates:
|
||
|
|
for _ in range(3):
|
||
|
|
asset = random.choice(ASSETS)
|
||
|
|
price = random.randint(100, 100000)
|
||
|
|
pct = random.randint(5, 50)
|
||
|
|
target = price + random.randint(100, 5000)
|
||
|
|
|
||
|
|
text = template.format(
|
||
|
|
asset=asset, price=price, pct=pct, target=target
|
||
|
|
)
|
||
|
|
|
||
|
|
labels = [0.0] * 6
|
||
|
|
labels[CFG_EMO.emotion_labels.index(emotion)] = 1.0
|
||
|
|
|
||
|
|
data.append({
|
||
|
|
"text": text,
|
||
|
|
"labels": labels,
|
||
|
|
"primary_emotion": emotion,
|
||
|
|
})
|
||
|
|
|
||
|
|
# Mixed emotion samples
|
||
|
|
mixed_templates = [
|
||
|
|
("greed,fear", "FOMO buying {asset} at ${price} but scared of reversal"),
|
||
|
|
("joy,fear", "Made 10x on {asset} at ${price} but terrified of giving back gains"),
|
||
|
|
("greed,anger", "Want to buy {asset} at ${price} but angry at manipulation"),
|
||
|
|
("fear,sadness", "Watching {asset} crash to ${price}, hopeless"),
|
||
|
|
("joy,greed", "To the moon with {asset} at ${price}! Loading more!"),
|
||
|
|
]
|
||
|
|
|
||
|
|
for (emos, template) in mixed_templates:
|
||
|
|
for _ in range(5):
|
||
|
|
asset = random.choice(ASSETS)
|
||
|
|
price = random.randint(100, 100000)
|
||
|
|
text = template.format(asset=asset, price=price)
|
||
|
|
labels = [0.0] * 6
|
||
|
|
for emo in emos.split(","):
|
||
|
|
labels[CFG_EMO.emotion_labels.index(emo)] = 1.0
|
||
|
|
data.append({"text": text, "labels": labels, "primary_emotion": emos})
|
||
|
|
|
||
|
|
print(f"🎭 Built {len(data)} emotion samples")
|
||
|
|
return data
|
||
|
|
|
||
|
|
# ============================================================
|
||
|
|
# MAIN TRAINING
|
||
|
|
# ============================================================
|
||
|
|
def train_emotion_lora():
|
||
|
|
if not TORCH_AVAILABLE:
|
||
|
|
print("❌ torch/peft not available")
|
||
|
|
return ""
|
||
|
|
|
||
|
|
print("="*60)
|
||
|
|
print("LoRA CRYPTO EMOTION PIPELINE")
|
||
|
|
print("="*60)
|
||
|
|
|
||
|
|
raw_data = build_emotion_dataset()
|
||
|
|
random.shuffle(raw_data)
|
||
|
|
|
||
|
|
split = int(0.9 * len(raw_data))
|
||
|
|
train_data = raw_data[:split]
|
||
|
|
val_data = raw_data[split:]
|
||
|
|
print(f"Train: {len(train_data)} | Val: {len(val_data)}")
|
||
|
|
|
||
|
|
label2id = {l: i for i, l in enumerate(CFG_EMO.emotion_labels)}
|
||
|
|
id2label = {i: l for i, l in enumerate(CFG_EMO.emotion_labels)}
|
||
|
|
|
||
|
|
def to_dataset(data):
|
||
|
|
return Dataset.from_dict({
|
||
|
|
"text": [d["text"] for d in data],
|
||
|
|
"labels": [d["labels"] for d in data],
|
||
|
|
})
|
||
|
|
|
||
|
|
train_ds = to_dataset(train_data)
|
||
|
|
val_ds = to_dataset(val_data)
|
||
|
|
|
||
|
|
tokenizer = AutoTokenizer.from_pretrained(CFG_EMO.base_model)
|
||
|
|
|
||
|
|
def tokenize(batch):
|
||
|
|
return tokenizer(batch["text"], truncation=True, max_length=CFG_EMO.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", "labels"])
|
||
|
|
val_ds.set_format("torch", columns=["input_ids", "attention_mask", "labels"])
|
||
|
|
|
||
|
|
# Model + LoRA - FIXED: proper from_pretrained call
|
||
|
|
model = AutoModelForSequenceClassification.from_pretrained(
|
||
|
|
CFG_EMO.base_model,
|
||
|
|
ignore_mismatched_sizes=True,
|
||
|
|
num_labels=6,
|
||
|
|
id2label=id2label,
|
||
|
|
label2id=label2id,
|
||
|
|
problem_type="multi_label_classification",
|
||
|
|
)
|
||
|
|
|
||
|
|
lora_config = LoraConfig(
|
||
|
|
r=CFG_EMO.r, lora_alpha=CFG_EMO.lora_alpha, lora_dropout=CFG_EMO.lora_dropout,
|
||
|
|
target_modules=list(CFG_EMO.target_modules), bias="none", task_type=TaskType.SEQ_CLS
|
||
|
|
)
|
||
|
|
model = get_peft_model(model, lora_config)
|
||
|
|
model.print_trainable_parameters()
|
||
|
|
|
||
|
|
# Weighted loss for crypto emotions
|
||
|
|
class_weights = torch.tensor([CFG_EMO.crypto_weights[l] for l in CFG_EMO.emotion_labels], dtype=torch.float)
|
||
|
|
|
||
|
|
class WeightedTrainer(Trainer):
|
||
|
|
def compute_loss(self, model, inputs, return_outputs=False, num_items_in_batch=None):
|
||
|
|
labels = inputs.pop("labels")
|
||
|
|
outputs = model(**inputs)
|
||
|
|
logits = outputs.logits
|
||
|
|
loss_fct = torch.nn.BCEWithLogitsLoss(pos_weight=class_weights.to(logits.device))
|
||
|
|
loss = loss_fct(logits, labels.float())
|
||
|
|
return (loss, outputs) if return_outputs else loss
|
||
|
|
|
||
|
|
output_dir = CFG_EMO.output_dir
|
||
|
|
os.makedirs(output_dir, exist_ok=True)
|
||
|
|
|
||
|
|
training_args = TrainingArguments(
|
||
|
|
output_dir=CFG_EMO.output_dir,
|
||
|
|
num_train_epochs=CFG_EMO.num_epochs,
|
||
|
|
per_device_train_batch_size=CFG_EMO.batch_size,
|
||
|
|
per_device_eval_batch_size=CFG_EMO.batch_size * 2,
|
||
|
|
gradient_accumulation_steps=CFG_EMO.grad_accum,
|
||
|
|
learning_rate=CFG_EMO.learning_rate,
|
||
|
|
warmup_ratio=CFG_EMO.warmup_ratio,
|
||
|
|
weight_decay=CFG_EMO.weight_decay,
|
||
|
|
max_grad_norm=1.0,
|
||
|
|
eval_strategy="steps",
|
||
|
|
eval_steps=CFG_EMO.save_steps,
|
||
|
|
save_strategy="steps",
|
||
|
|
save_steps=CFG_EMO.save_steps,
|
||
|
|
save_total_limit=CFG_EMO.save_total_limit,
|
||
|
|
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,
|
||
|
|
remove_unused_columns=False,
|
||
|
|
report_to="none",
|
||
|
|
seed=42,
|
||
|
|
)
|
||
|
|
|
||
|
|
def compute_metrics(eval_pred):
|
||
|
|
from sklearn.metrics import f1_score
|
||
|
|
logits, labels = eval_pred
|
||
|
|
preds = (torch.sigmoid(torch.tensor(logits)) > 0.5).int().numpy()
|
||
|
|
return {
|
||
|
|
"f1_macro": f1_score(labels, preds, average="macro", zero_division=0),
|
||
|
|
"f1_micro": f1_score(labels, preds, average="micro", zero_division=0),
|
||
|
|
"accuracy": (preds == labels).mean(),
|
||
|
|
}
|
||
|
|
|
||
|
|
trainer = WeightedTrainer(
|
||
|
|
model=model,
|
||
|
|
args=TrainingArguments(
|
||
|
|
output_dir=CFG_EMO.output_dir,
|
||
|
|
num_train_epochs=CFG_EMO.num_epochs,
|
||
|
|
per_device_train_batch_size=CFG_EMO.batch_size,
|
||
|
|
per_device_eval_batch_size=CFG_EMO.batch_size * 2,
|
||
|
|
gradient_accumulation_steps=CFG_EMO.grad_accum,
|
||
|
|
learning_rate=CFG_EMO.learning_rate,
|
||
|
|
warmup_ratio=CFG_EMO.warmup_ratio,
|
||
|
|
weight_decay=CFG_EMO.weight_decay,
|
||
|
|
max_grad_norm=1.0,
|
||
|
|
eval_strategy="steps",
|
||
|
|
eval_steps=CFG_EMO.save_steps,
|
||
|
|
save_strategy="steps",
|
||
|
|
save_steps=CFG_EMO.save_steps,
|
||
|
|
save_total_limit=CFG_EMO.save_total_limit,
|
||
|
|
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,
|
||
|
|
remove_unused_columns=False,
|
||
|
|
report_to="none",
|
||
|
|
seed=42,
|
||
|
|
),
|
||
|
|
train_dataset=train_ds,
|
||
|
|
eval_dataset=val_ds,
|
||
|
|
tokenizer=AutoTokenizer.from_pretrained(CFG_EMO.base_model),
|
||
|
|
compute_metrics=compute_metrics,
|
||
|
|
callbacks=[EarlyStoppingCallback(early_stopping_patience=3, early_stopping_threshold=0.001)],
|
||
|
|
)
|
||
|
|
|
||
|
|
# Prepare datasets
|
||
|
|
tokenizer = AutoTokenizer.from_pretrained(CFG_EMO.base_model)
|
||
|
|
|
||
|
|
def tokenize(batch):
|
||
|
|
return tokenizer(batch["text"], truncation=True, max_length=CFG_EMO.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", "labels"])
|
||
|
|
val_ds.set_format("torch", columns=["input_ids", "attention_mask", "labels"])
|
||
|
|
|
||
|
|
trainer = WeightedTrainer(
|
||
|
|
model=model,
|
||
|
|
args=training_args,
|
||
|
|
train_dataset=train_ds,
|
||
|
|
eval_dataset=val_ds,
|
||
|
|
tokenizer=AutoTokenizer.from_pretrained(CFG_EMO.base_model),
|
||
|
|
compute_metrics=compute_metrics,
|
||
|
|
callbacks=[EarlyStoppingCallback(early_stopping_patience=3, early_stopping_threshold=0.001)],
|
||
|
|
)
|
||
|
|
|
||
|
|
print("🏋️ Training emotion LoRA...")
|
||
|
|
trainer.train()
|
||
|
|
|
||
|
|
best_path = os.path.join(CFG_EMO.output_dir, "best")
|
||
|
|
trainer.save_model(best_path)
|
||
|
|
AutoTokenizer.from_pretrained(CFG_EMO.base_model).save_pretrained(best_path)
|
||
|
|
|
||
|
|
# Cleanup checkpoints
|
||
|
|
for d in os.listdir(CFG_EMO.output_dir):
|
||
|
|
if d.startswith("checkpoint-"):
|
||
|
|
shutil.rmtree(os.path.join(CFG_EMO.output_dir, d), ignore_errors=True)
|
||
|
|
|
||
|
|
print(f"✅ Best emotion model saved to: {best_path}")
|
||
|
|
return best_path
|
||
|
|
|
||
|
|
if __name__ == "__main__":
|
||
|
|
import argparse
|
||
|
|
import os
|
||
|
|
import shutil
|
||
|
|
from transformers import AutoTokenizer, AutoModelForSequenceClassification, TrainingArguments, Trainer, EarlyStoppingCallback
|
||
|
|
from peft import LoraConfig, get_peft_model, TaskType
|
||
|
|
from datasets import Dataset
|
||
|
|
from sklearn.metrics import f1_score
|
||
|
|
|
||
|
|
parser = argparse.ArgumentParser()
|
||
|
|
parser.add_argument("--epochs", type=int, default=5)
|
||
|
|
parser.add_argument("--batch", type=int, default=8)
|
||
|
|
parser.add_argument("--r", type=int, default=8)
|
||
|
|
args = parser.parse_args()
|
||
|
|
|
||
|
|
CFG_EMO.num_epochs = args.epochs
|
||
|
|
CFG_EMO.batch_size = args.batch
|
||
|
|
CFG_EMO.r = args.r
|
||
|
|
|
||
|
|
train_emotion_lora()
|