""" Character-level embeddings using CNN or BiLSTM Part of BioBBC enhanced features """ import torch import torch.nn as nn class CharCNN(nn.Module): """ Character-level CNN embeddings Captures morphological features from character sequences """ def __init__(self, char_vocab_size, char_embedding_dim=50, char_hidden_dim=50, kernel_sizes=[3, 4, 5], dropout=0.3): """ Args: char_vocab_size: Size of character vocabulary char_embedding_dim: Dimension of character embeddings char_hidden_dim: Hidden dimension for convolution output kernel_sizes: List of kernel sizes for multi-scale CNN dropout: Dropout rate """ super(CharCNN, self).__init__() self.char_embedding = nn.Embedding(char_vocab_size, char_embedding_dim, padding_idx=0) # Multiple CNN layers with different kernel sizes self.convs = nn.ModuleList([ nn.Conv1d( in_channels=char_embedding_dim, out_channels=char_hidden_dim, kernel_size=k ) for k in kernel_sizes ]) self.dropout = nn.Dropout(dropout) self.output_dim = char_hidden_dim * len(kernel_sizes) def forward(self, char_ids): """ Args: char_ids: (batch_size, seq_len, max_word_len) - character IDs Returns: Character embeddings: (batch_size, seq_len, output_dim) """ batch_size, seq_len, max_word_len = char_ids.size() # Reshape to (batch_size * seq_len, max_word_len) char_ids = char_ids.view(batch_size * seq_len, max_word_len) # Embed characters: (batch_size * seq_len, max_word_len, char_embedding_dim) char_embeds = self.char_embedding(char_ids) # Transpose for Conv1d: (batch_size * seq_len, char_embedding_dim, max_word_len) char_embeds = char_embeds.transpose(1, 2) # Apply convolutions and max pooling conv_outputs = [] for conv in self.convs: # Conv: (batch_size * seq_len, char_hidden_dim, seq_len') conv_out = torch.relu(conv(char_embeds)) # Max pool: (batch_size * seq_len, char_hidden_dim) pooled = torch.max(conv_out, dim=2)[0] conv_outputs.append(pooled) # Concatenate outputs from different kernel sizes char_features = torch.cat(conv_outputs, dim=1) # (batch_size * seq_len, output_dim) # Reshape back: (batch_size, seq_len, output_dim) char_features = char_features.view(batch_size, seq_len, -1) char_features = self.dropout(char_features) return char_features class CharBiLSTM(nn.Module): """ Character-level BiLSTM embeddings Alternative to CNN for character representation """ def __init__(self, char_vocab_size, char_embedding_dim=50, char_hidden_dim=50, dropout=0.3): """ Args: char_vocab_size: Size of character vocabulary char_embedding_dim: Dimension of character embeddings char_hidden_dim: Hidden dimension for BiLSTM dropout: Dropout rate """ super(CharBiLSTM, self).__init__() self.char_embedding = nn.Embedding(char_vocab_size, char_embedding_dim, padding_idx=0) self.char_lstm = nn.LSTM( input_size=char_embedding_dim, hidden_size=char_hidden_dim, num_layers=1, batch_first=True, bidirectional=True ) self.dropout = nn.Dropout(dropout) self.output_dim = char_hidden_dim * 2 # Bidirectional def forward(self, char_ids): """ Args: char_ids: (batch_size, seq_len, max_word_len) - character IDs Returns: Character embeddings: (batch_size, seq_len, output_dim) """ batch_size, seq_len, max_word_len = char_ids.size() # Reshape to (batch_size * seq_len, max_word_len) char_ids = char_ids.view(batch_size * seq_len, max_word_len) # Embed characters: (batch_size * seq_len, max_word_len, char_embedding_dim) char_embeds = self.char_embedding(char_ids) # BiLSTM: (batch_size * seq_len, max_word_len, char_hidden_dim * 2) lstm_out, (h_n, c_n) = self.char_lstm(char_embeds) # Take last hidden states from both directions # h_n: (2, batch_size * seq_len, char_hidden_dim) forward_hidden = h_n[0] # (batch_size * seq_len, char_hidden_dim) backward_hidden = h_n[1] # (batch_size * seq_len, char_hidden_dim) # Concatenate: (batch_size * seq_len, char_hidden_dim * 2) char_features = torch.cat([forward_hidden, backward_hidden], dim=1) # Reshape back: (batch_size, seq_len, output_dim) char_features = char_features.view(batch_size, seq_len, -1) char_features = self.dropout(char_features) return char_features class CharacterVocab: """Character vocabulary builder""" def __init__(self, texts=None): """ Args: texts: List of text strings to build vocabulary from """ self.char2id = {'': 0, '': 1} self.id2char = {0: '', 1: ''} if texts: self.build_vocab(texts) def build_vocab(self, texts): """Build character vocabulary from texts""" chars = set() for text in texts: chars.update(list(text)) # Sort for consistency for char in sorted(chars): if char not in self.char2id: idx = len(self.char2id) self.char2id[char] = idx self.id2char[idx] = char def encode(self, word, max_len=20): """ Encode word to character IDs Args: word: Word string max_len: Maximum word length Returns: List of character IDs """ char_ids = [self.char2id.get(c, self.char2id['']) for c in word[:max_len]] # Pad if needed if len(char_ids) < max_len: char_ids += [self.char2id['']] * (max_len - len(char_ids)) return char_ids def __len__(self): return len(self.char2id)