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]