""" Enhanced BioBBC Model: BERT + Character + POS + Domain Embeddings + BiLSTM + CRF Implements the full BioBBC architecture with all enhanced features """ import torch import torch.nn as nn from transformers import AutoModel from crf_layer import CRF from char_embeddings import CharCNN, CharBiLSTM from pos_embeddings import POSEmbedding from domain_embeddings import DomainEmbedding from config import Config class BioBBC_Model(nn.Module): """ BioBBC: Enhanced BERT-BiLSTM-CRF with multiple feature types Features concatenated: 1. BERT embeddings (contextual) 2. Character embeddings (morphological) 3. POS embeddings (syntactic) 4. Domain-specific embeddings (biomedical knowledge) """ def __init__(self, bert_model_name, hidden_dim, num_layers, num_labels, char_vocab_size=None, pos_vocab_size=None, domain_vocab=None, use_char_cnn=True, use_char=True, use_pos=True, use_domain=True, char_embedding_dim=50, char_hidden_dim=50, pos_embedding_dim=25, domain_embedding_dim=200, domain_embedding_path=None, dropout=0.3): """ Args: bert_model_name: Pre-trained BERT model hidden_dim: BiLSTM hidden dimension num_layers: BiLSTM layers num_labels: Number of NER labels char_vocab_size: Character vocabulary size pos_vocab_size: POS vocabulary size domain_vocab: Domain word vocabulary use_char_cnn: Use CNN (True) or BiLSTM (False) for characters use_char: Enable character embeddings use_pos: Enable POS embeddings use_domain: Enable domain embeddings char_embedding_dim: Character embedding dimension char_hidden_dim: Character hidden dimension pos_embedding_dim: POS embedding dimension domain_embedding_dim: Domain embedding dimension domain_embedding_path: Path to pre-trained domain embeddings dropout: Dropout rate """ super(BioBBC_Model, self).__init__() # BERT encoder self.bert = AutoModel.from_pretrained(bert_model_name) self.bert_hidden_size = self.bert.config.hidden_size # Feature flags self.use_char = use_char and char_vocab_size is not None self.use_pos = use_pos and pos_vocab_size is not None self.use_domain = use_domain and domain_vocab is not None # Calculate total input dimension total_input_dim = self.bert_hidden_size # Character embeddings if self.use_char: if use_char_cnn: self.char_encoder = CharCNN( char_vocab_size=char_vocab_size, char_embedding_dim=char_embedding_dim, char_hidden_dim=char_hidden_dim, dropout=dropout ) else: self.char_encoder = CharBiLSTM( char_vocab_size=char_vocab_size, char_embedding_dim=char_embedding_dim, char_hidden_dim=char_hidden_dim, dropout=dropout ) total_input_dim += self.char_encoder.output_dim # POS embeddings if self.use_pos: self.pos_encoder = POSEmbedding( pos_vocab_size=pos_vocab_size, pos_embedding_dim=pos_embedding_dim, dropout=dropout ) total_input_dim += self.pos_encoder.output_dim # Domain-specific embeddings if self.use_domain: self.domain_encoder = DomainEmbedding( embedding_path=domain_embedding_path, embedding_dim=domain_embedding_dim, vocab=domain_vocab, freeze=True, # Typically freeze pre-trained embeddings dropout=dropout ) total_input_dim += self.domain_encoder.output_dim # BiLSTM layer (processes concatenated features) self.bilstm = nn.LSTM( input_size=total_input_dim, hidden_size=hidden_dim, num_layers=num_layers, batch_first=True, bidirectional=True, dropout=dropout if num_layers > 1 else 0 ) self.dropout = nn.Dropout(dropout) # Linear projection to label space self.hidden2label = nn.Linear(hidden_dim * 2, num_labels) # CRF layer self.crf = CRF(num_labels, batch_first=True) self.num_labels = num_labels self.total_input_dim = total_input_dim print(f"\nBioBBC Model Initialized:") print(f" BERT dimension: {self.bert_hidden_size}") print(f" Character embeddings: {self.use_char} (dim: {self.char_encoder.output_dim if self.use_char else 0})") print(f" POS embeddings: {self.use_pos} (dim: {self.pos_encoder.output_dim if self.use_pos else 0})") print(f" Domain embeddings: {self.use_domain} (dim: {self.domain_encoder.output_dim if self.use_domain else 0})") print(f" Total input dimension: {total_input_dim}") print(f" BiLSTM hidden dimension: {hidden_dim} x 2 (bidirectional)") def forward(self, input_ids, attention_mask, char_ids=None, pos_ids=None, domain_ids=None, labels=None): """ Forward pass with multiple feature types Args: input_ids: (batch_size, seq_len) - BERT input IDs attention_mask: (batch_size, seq_len) - attention mask char_ids: (batch_size, seq_len, max_word_len) - character IDs (optional) pos_ids: (batch_size, seq_len) - POS tag IDs (optional) domain_ids: (batch_size, seq_len) - domain word IDs (optional) labels: (batch_size, seq_len) - gold labels (optional, for training) Returns: If labels provided: loss Else: predicted tag sequences """ # 1. BERT encoding bert_outputs = self.bert( input_ids=input_ids, attention_mask=attention_mask ) bert_embeds = bert_outputs.last_hidden_state # (batch_size, seq_len, bert_hidden_size) # 2. Concatenate additional features feature_list = [bert_embeds] # Character embeddings if self.use_char and char_ids is not None: char_embeds = self.char_encoder(char_ids) # (batch_size, seq_len, char_dim) feature_list.append(char_embeds) # POS embeddings if self.use_pos and pos_ids is not None: pos_embeds = self.pos_encoder(pos_ids) # (batch_size, seq_len, pos_dim) feature_list.append(pos_embeds) # Domain embeddings if self.use_domain and domain_ids is not None: domain_embeds = self.domain_encoder(domain_ids) # (batch_size, seq_len, domain_dim) feature_list.append(domain_embeds) # Concatenate all features combined_features = torch.cat(feature_list, dim=2) # (batch_size, seq_len, total_dim) # 3. BiLSTM lstm_output, _ = self.bilstm(combined_features) # (batch_size, seq_len, hidden_dim*2) lstm_output = self.dropout(lstm_output) # 4. Project to label space emissions = self.hidden2label(lstm_output) # (batch_size, seq_len, num_labels) # 5. CRF if labels is not None: # Training mode: compute loss # Replace -100 (ignore index) with 0 for CRF # Create mask that excludes positions with -100 labels labels_for_crf = labels.clone() mask = (labels != -100) & attention_mask.bool() labels_for_crf[labels_for_crf == -100] = 0 loss = self.crf.compute_loss(emissions, labels_for_crf, mask) return loss else: # Inference mode: decode best sequence mask = attention_mask.bool() predictions = self.crf(emissions, mask) return predictions def get_feature_importance(self, input_ids, attention_mask, char_ids=None, pos_ids=None, domain_ids=None): """ Analyze contribution of each feature type Returns embeddings for each feature separately """ with torch.no_grad(): # BERT bert_outputs = self.bert(input_ids=input_ids, attention_mask=attention_mask) bert_embeds = bert_outputs.last_hidden_state features = {'bert': bert_embeds} # Character if self.use_char and char_ids is not None: features['char'] = self.char_encoder(char_ids) # POS if self.use_pos and pos_ids is not None: features['pos'] = self.pos_encoder(pos_ids) # Domain if self.use_domain and domain_ids is not None: features['domain'] = self.domain_encoder(domain_ids) return features def freeze_bert(self): """Freeze BERT parameters""" for param in self.bert.parameters(): param.requires_grad = False def unfreeze_bert(self): """Unfreeze BERT parameters""" for param in self.bert.parameters(): param.requires_grad = True def create_biobbc_model(char_vocab_size, pos_vocab_size, domain_vocab, config=None): """ Factory function to create BioBBC model Args: char_vocab_size: Character vocabulary size pos_vocab_size: POS vocabulary size domain_vocab: Domain word vocabulary dict config: Configuration object Returns: BioBBC_Model instance """ if config is None: config = Config model = BioBBC_Model( bert_model_name=config.BERT_MODEL, hidden_dim=config.HIDDEN_DIM, num_layers=config.NUM_LAYERS, num_labels=config.NUM_LABELS, char_vocab_size=char_vocab_size, pos_vocab_size=pos_vocab_size, domain_vocab=domain_vocab, use_char_cnn=True, use_char=config.USE_CHAR_EMBEDDINGS, use_pos=config.USE_POS_EMBEDDINGS, use_domain=config.USE_DOMAIN_EMBEDDINGS, char_embedding_dim=config.CHAR_EMBEDDING_DIM, char_hidden_dim=config.CHAR_HIDDEN_DIM, pos_embedding_dim=config.POS_EMBEDDING_DIM, domain_embedding_dim=config.DOMAIN_EMBEDDING_DIM, domain_embedding_path=config.DOMAIN_EMBEDDING_PATH if os.path.exists(config.DOMAIN_EMBEDDING_PATH) else None, dropout=config.DROPOUT ) return model.to(config.DEVICE) import os