""" POS (Part-of-Speech) tag embeddings using spaCy Part of BioBBC enhanced features - syntactic information """ import torch import torch.nn as nn import spacy class POSEmbedding(nn.Module): """ POS tag embeddings layer Provides syntactic features to complement BERT """ def __init__(self, pos_vocab_size, pos_embedding_dim=25, dropout=0.3): """ Args: pos_vocab_size: Size of POS tag vocabulary pos_embedding_dim: Dimension of POS embeddings dropout: Dropout rate """ super(POSEmbedding, self).__init__() self.pos_embedding = nn.Embedding(pos_vocab_size, pos_embedding_dim, padding_idx=0) self.dropout = nn.Dropout(dropout) self.output_dim = pos_embedding_dim def forward(self, pos_ids): """ Args: pos_ids: (batch_size, seq_len) - POS tag IDs Returns: POS embeddings: (batch_size, seq_len, pos_embedding_dim) """ pos_embeds = self.pos_embedding(pos_ids) pos_embeds = self.dropout(pos_embeds) return pos_embeds class POSTagger: """ POS tagger using spaCy Extracts POS tags for biomedical text """ def __init__(self, model_name='en_core_web_sm'): """ Args: model_name: spaCy model name For biomedical: 'en_core_sci_sm' (scispacy) For general: 'en_core_web_sm' """ try: self.nlp = spacy.load(model_name) except OSError: print(f"spaCy model '{model_name}' not found. Downloading...") import subprocess subprocess.run(['python', '-m', 'spacy', 'download', model_name]) self.nlp = spacy.load(model_name) # Disable unnecessary components for speed self.nlp.disable_pipes(['parser', 'ner']) # Build POS vocabulary self.pos2id = {'': 0, '': 1} self.id2pos = {0: '', 1: ''} self._build_pos_vocab() def _build_pos_vocab(self): """Build POS tag vocabulary from spaCy""" # Common POS tags (Universal POS tags) common_pos = [ 'ADJ', 'ADP', 'ADV', 'AUX', 'CCONJ', 'DET', 'INTJ', 'NOUN', 'NUM', 'PART', 'PRON', 'PROPN', 'PUNCT', 'SCONJ', 'SYM', 'VERB', 'X', 'SPACE' ] for pos in common_pos: if pos not in self.pos2id: idx = len(self.pos2id) self.pos2id[pos] = idx self.id2pos[idx] = pos def tag_tokens(self, tokens): """ Extract POS tags for list of tokens Args: tokens: List of token strings Returns: List of POS tag IDs """ # Join tokens and process with spaCy text = ' '.join(tokens) doc = self.nlp(text) # Extract POS tags pos_tags = [] doc_tokens = [token.text for token in doc] # Align with input tokens token_idx = 0 for token in tokens: if token_idx < len(doc_tokens) and token == doc_tokens[token_idx]: pos_tag = doc[token_idx].pos_ pos_id = self.pos2id.get(pos_tag, self.pos2id['']) pos_tags.append(pos_id) token_idx += 1 else: # If alignment fails, use UNK pos_tags.append(self.pos2id['']) return pos_tags def tag_batch(self, batch_tokens): """ Extract POS tags for batch of token lists Args: batch_tokens: List of token lists Returns: List of POS tag ID lists """ return [self.tag_tokens(tokens) for tokens in batch_tokens] def get_vocab_size(self): """Return POS vocabulary size""" return len(self.pos2id) class SciBioPOSTagger(POSTagger): """ Biomedical-specific POS tagger using scispaCy Better for biomedical text """ def __init__(self): """Initialize with biomedical spaCy model""" try: import scispacy super().__init__(model_name='en_core_sci_sm') except ImportError: print("scispaCy not installed. Installing...") import subprocess subprocess.run(['pip', 'install', 'scispacy']) subprocess.run(['pip', 'install', 'https://s3-us-west-2.amazonaws.com/ai2-s2-scispacy/releases/v0.5.1/en_core_sci_sm-0.5.1.tar.gz']) import scispacy super().__init__(model_name='en_core_sci_sm') def get_pos_tagger(use_biomedical=True): """ Factory function to get appropriate POS tagger Args: use_biomedical: If True, use biomedical-specific tagger Returns: POSTagger instance """ if use_biomedical: try: return SciBioPOSTagger() except Exception as e: print(f"Failed to load biomedical tagger: {e}") print("Falling back to general tagger...") return POSTagger() else: return POSTagger()