stockproject / brain /analysis /sentiment.py
harshisageek's picture
deploy: clean history for HuggingFace
1cd56b6
Raw History Blame Contribute Delete
7.85 kB
import torch
import torch.nn as nn
import logging
import os
from transformers import DistilBertTokenizer, DistilBertModel
from typing import List
from threading import Lock
from brain.core.exceptions import ModelLoadException, AnalysisException
logger = logging.getLogger(__name__)
# Path to custom fine-tuned DistilBERT model
PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
CUSTOM_MODEL_DIR = os.path.join(PROJECT_ROOT, "models", "custom_sentiment")
class FinancialSentimentModel(nn.Module):
"""
Custom DistilBERT + classification head — fine-tuned on Financial PhraseBank.
Architecture: DistilBERT(768) → Dropout(0.3) → Dense(256) → ReLU → Dropout(0.15) → Dense(3)
"""
def __init__(self, model_name="distilbert-base-uncased", num_classes=3, dropout=0.3):
super(FinancialSentimentModel, self).__init__()
self.distilbert = DistilBertModel.from_pretrained(model_name)
self.classifier = nn.Sequential(
nn.Dropout(dropout),
nn.Linear(self.distilbert.config.hidden_size, 256),
nn.ReLU(),
nn.Dropout(dropout / 2),
nn.Linear(256, num_classes)
)
def forward(self, input_ids, attention_mask):
outputs = self.distilbert(input_ids=input_ids, attention_mask=attention_mask)
cls_output = outputs.last_hidden_state[:, 0, :]
return self.classifier(cls_output)
class SentimentEngine:
"""
NLP Engine using custom fine-tuned DistilBERT for sentiment analysis.
Falls back to FinBERT if custom model not available.
Thread-safe Singleton implementation.
"""
_model = None
_tokenizer = None
_device = None
_is_finbert_fallback = False
_lock = Lock()
@classmethod
def _load_model(cls):
if cls._model is not None:
return
with cls._lock:
if cls._model is not None:
return
model_path = os.path.join(CUSTOM_MODEL_DIR, "sentiment_model.pt")
if os.path.exists(model_path):
logger.info("Initializing CUSTOM fine-tuned DistilBERT model...")
try:
if torch.cuda.is_available():
cls._device = torch.device("cuda")
logger.info(f"Custom DistilBERT: GPU {torch.cuda.get_device_name(0)}")
elif hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
cls._device = torch.device("mps")
logger.info("Custom DistilBERT: MPS")
else:
cls._device = torch.device("cpu")
logger.info("Custom DistilBERT: CPU")
cls._tokenizer = DistilBertTokenizer.from_pretrained(CUSTOM_MODEL_DIR)
cls._model = FinancialSentimentModel("distilbert-base-uncased")
checkpoint = torch.load(model_path, map_location=cls._device, weights_only=False)
sd = checkpoint["model_state_dict"]
sd_fp32 = {k: v.float() if v.dtype == torch.float16 else v for k, v in sd.items()}
cls._model.load_state_dict(sd_fp32)
cls._model.to(cls._device)
cls._model.eval()
cls._is_finbert_fallback = False
val_acc = checkpoint.get("val_acc", "N/A")
logger.info(f"Custom DistilBERT loaded (val_acc: {val_acc})")
return
except Exception as e:
logger.warning(f"Custom DistilBERT load failed: {e}. Falling back.")
# Fallback to FinBERT
logger.info("Custom model not found — falling back to FinBERT...")
try:
from transformers import AutoTokenizer, AutoModelForSequenceClassification
fin_tokenizer = AutoTokenizer.from_pretrained("yiyanghkust/finbert-tone")
fin_model = AutoModelForSequenceClassification.from_pretrained("yiyanghkust/finbert-tone")
fin_model.eval()
cls._device = torch.device("cpu")
cls._model = fin_model
cls._tokenizer = fin_tokenizer
cls._is_finbert_fallback = True
logger.info("FinBERT fallback initialized.")
except Exception as e:
logger.critical(f"FinBERT Load Failed: {e}")
raise ModelLoadException(f"Could not load any sentiment model: {e}")
@classmethod
def analyze_batch(cls, texts: List[str]) -> List[float]:
"""
Analyzes a batch of texts.
Returns sentiment scores from -1.0 (Negative) to 1.0 (Positive).
"""
if not texts:
return []
cls._load_model()
# Smart Truncation
cleaned_texts = []
for t in texts:
if not t:
cleaned_texts.append("")
continue
if len(t) > 1500:
cleaned_texts.append(t[:1000] + " ... " + t[-500:])
else:
cleaned_texts.append(t)
valid_indices = [i for i, t in enumerate(cleaned_texts) if t.strip()]
valid_inputs = [cleaned_texts[i] for i in valid_indices]
if not valid_inputs:
return [0.0] * len(texts)
try:
if cls._is_finbert_fallback:
return cls._finbert_inference(valid_inputs, valid_indices, len(texts))
# --- Custom DistilBERT inference ---
final_scores = [0.0] * len(texts)
batch_size = 16
with torch.no_grad():
for i in range(0, len(valid_inputs), batch_size):
batch_texts = valid_inputs[i:i + batch_size]
batch_indices = valid_indices[i:i + batch_size]
encoding = cls._tokenizer(
batch_texts,
max_length=128,
padding="max_length",
truncation=True,
return_tensors="pt"
).to(cls._device)
logits = cls._model(encoding["input_ids"], encoding["attention_mask"])
probs = torch.softmax(logits, dim=1)
for j, idx in enumerate(batch_indices):
pos_prob = probs[j, 2].item()
neg_prob = probs[j, 0].item()
composite = (pos_prob - neg_prob) * 0.95
final_scores[idx] = round(composite, 4)
return final_scores
except Exception as e:
logger.error(f"Sentiment Batch Error: {e}")
raise AnalysisException(f"Sentiment analysis failed: {e}")
@classmethod
def _finbert_inference(cls, valid_inputs, valid_indices, total_len):
"""Handle FinBERT fallback."""
# FinBERT id2label: {0: Neutral, 1: Positive, 2: Negative}
final_scores = [0.0] * total_len
with torch.no_grad():
for i in range(0, len(valid_inputs), 16):
batch = valid_inputs[i:i+16]
enc = cls._tokenizer(batch, max_length=128, padding="max_length",
truncation=True, return_tensors="pt")
outputs = cls._model(**enc)
probs = torch.softmax(outputs.logits, dim=1)
for j, idx in enumerate(valid_indices[i:i+16]):
pos = probs[j, 1].item()
neg = probs[j, 2].item()
composite = (pos - neg) * 0.95
final_scores[idx] = round(composite, 4)
return final_scores
@classmethod
def analyze_one(cls, text: str) -> float:
return cls.analyze_batch([text])[0]