""" Domain-specific word embeddings for biomedical text Supports BioWordVec, PubMed embeddings, and other pre-trained biomedical embeddings Part of BioBBC enhanced features """ import torch import torch.nn as nn import numpy as np import os class DomainEmbedding(nn.Module): """ Domain-specific word embeddings layer Loads pre-trained biomedical word embeddings """ def __init__(self, embedding_path, embedding_dim=200, vocab=None, freeze=False, dropout=0.3): """ Args: embedding_path: Path to pre-trained embeddings file embedding_dim: Dimension of embeddings vocab: Word vocabulary (word2id dict) freeze: If True, embeddings are not updated during training dropout: Dropout rate """ super(DomainEmbedding, self).__init__() self.embedding_dim = embedding_dim self.vocab = vocab # Initialize embedding layer if vocab is not None: vocab_size = len(vocab) else: vocab_size = 10000 # Default size, will be updated when loading self.word_embedding = nn.Embedding(vocab_size, embedding_dim, padding_idx=0) # Load pre-trained embeddings if path exists if embedding_path and os.path.exists(embedding_path): self._load_embeddings(embedding_path) if freeze: self.word_embedding.weight.requires_grad = False self.dropout = nn.Dropout(dropout) self.output_dim = embedding_dim def _load_embeddings(self, embedding_path): """ Load pre-trained embeddings from file Supports Word2Vec text format and GloVe format """ print(f"Loading domain embeddings from {embedding_path}...") embeddings_index = {} with open(embedding_path, 'r', encoding='utf-8') as f: for i, line in enumerate(f): if i == 0: # Skip header if exists (Word2Vec format) parts = line.strip().split() if len(parts) == 2: continue values = line.strip().split() if len(values) < 10: # Skip malformed lines continue word = values[0] try: vector = np.asarray(values[1:], dtype='float32') if len(vector) == self.embedding_dim: embeddings_index[word] = vector except ValueError: continue print(f"Loaded {len(embeddings_index)} word vectors") # Create embedding matrix if self.vocab: embedding_matrix = np.zeros((len(self.vocab), self.embedding_dim)) found = 0 for word, idx in self.vocab.items(): embedding_vector = embeddings_index.get(word.lower()) if embedding_vector is not None: embedding_matrix[idx] = embedding_vector found += 1 else: # Random initialization for OOV words embedding_matrix[idx] = np.random.normal(0, 0.1, self.embedding_dim) print(f"Matched {found}/{len(self.vocab)} words in vocabulary") # Set weights self.word_embedding.weight.data.copy_(torch.from_numpy(embedding_matrix)) def forward(self, word_ids): """ Args: word_ids: (batch_size, seq_len) - word IDs Returns: Domain embeddings: (batch_size, seq_len, embedding_dim) """ embeds = self.word_embedding(word_ids) embeds = self.dropout(embeds) return embeds class BioWordVecEmbedding(DomainEmbedding): """ BioWordVec embeddings specifically Pre-trained on biomedical text (PubMed + MIMIC-III) Download from: https://github.com/ncbi-nlp/BioWordVec """ def __init__(self, embedding_path='embeddings/BioWordVec_PubMed_MIMICIII_d200.txt', vocab=None, freeze=False, dropout=0.3): super().__init__( embedding_path=embedding_path, embedding_dim=200, vocab=vocab, freeze=freeze, dropout=dropout ) class PubMedEmbedding(DomainEmbedding): """ PubMed Word2Vec embeddings Pre-trained on PubMed abstracts Download from: http://evexdb.org/pmresources/vec-space-models/ """ def __init__(self, embedding_path='embeddings/PubMed_w2v.txt', vocab=None, freeze=False, dropout=0.3): super().__init__( embedding_path=embedding_path, embedding_dim=200, vocab=vocab, freeze=freeze, dropout=dropout ) class DomainVocab: """Build vocabulary for domain embeddings""" def __init__(self, texts=None, min_freq=1): """ Args: texts: List of token lists to build vocabulary from min_freq: Minimum frequency to include word """ self.word2id = {'': 0, '': 1} self.id2word = {0: '', 1: ''} self.word_freq = {} if texts: self.build_vocab(texts, min_freq) def build_vocab(self, texts, min_freq=1): """Build vocabulary from texts""" # Count frequencies for tokens in texts: for token in tokens: self.word_freq[token] = self.word_freq.get(token, 0) + 1 # Add words above min_freq for word, freq in sorted(self.word_freq.items()): if freq >= min_freq and word not in self.word2id: idx = len(self.word2id) self.word2id[word] = idx self.id2word[idx] = word def encode(self, tokens): """ Encode tokens to word IDs Args: tokens: List of token strings Returns: List of word IDs """ return [self.word2id.get(token, self.word2id['']) for token in tokens] def __len__(self): return len(self.word2id) def __getitem__(self, key): if isinstance(key, str): return self.word2id.get(key, self.word2id['']) else: return self.id2word.get(key, '') def load_biomedical_embeddings(embedding_type='biowordvec', embedding_path=None, vocab=None, freeze=True): """ Factory function to load biomedical embeddings Args: embedding_type: 'biowordvec', 'pubmed', or 'custom' embedding_path: Path to embedding file (for custom) vocab: Word vocabulary freeze: If True, don't update embeddings during training Returns: DomainEmbedding instance """ if embedding_type.lower() == 'biowordvec': return BioWordVecEmbedding(vocab=vocab, freeze=freeze) elif embedding_type.lower() == 'pubmed': return PubMedEmbedding(vocab=vocab, freeze=freeze) elif embedding_type.lower() == 'custom' and embedding_path: return DomainEmbedding(embedding_path, vocab=vocab, freeze=freeze) else: # No pre-trained embeddings, use random initialization print("No pre-trained embeddings found, using random initialization") return DomainEmbedding(embedding_path=None, vocab=vocab, freeze=False)