""" Conditional Random Field (CRF) layer for sequence labeling Implements CRF for enforcing label constraints """ import torch import torch.nn as nn class CRF(nn.Module): """Conditional Random Field layer""" def __init__(self, num_labels, batch_first=True): """ Args: num_labels: Number of labels/tags batch_first: If True, batch dimension is first """ super(CRF, self).__init__() self.num_labels = num_labels self.batch_first = batch_first # Transition parameters: transitions[i, j] = score of transitioning from j to i self.transitions = nn.Parameter(torch.randn(num_labels, num_labels)) # Start and end transitions self.start_transitions = nn.Parameter(torch.randn(num_labels)) self.end_transitions = nn.Parameter(torch.randn(num_labels)) self._initialize_parameters() def _initialize_parameters(self): """Initialize transition parameters""" nn.init.uniform_(self.transitions, -0.1, 0.1) nn.init.uniform_(self.start_transitions, -0.1, 0.1) nn.init.uniform_(self.end_transitions, -0.1, 0.1) def forward(self, emissions, mask=None): """ Decode best tag sequence using Viterbi algorithm Args: emissions: (batch_size, seq_len, num_labels) - emission scores mask: (batch_size, seq_len) - mask for valid positions Returns: List of best tag sequences """ if self.batch_first: emissions = emissions.transpose(0, 1) # (seq_len, batch_size, num_labels) if mask is not None: mask = mask.transpose(0, 1) # (seq_len, batch_size) return self._viterbi_decode(emissions, mask) def _viterbi_decode(self, emissions, mask=None): """Viterbi algorithm for decoding""" seq_len, batch_size, num_labels = emissions.shape if mask is None: mask = torch.ones(seq_len, batch_size, dtype=torch.bool, device=emissions.device) # Initialize score = self.start_transitions + emissions[0] # (batch_size, num_labels) history = [] # Forward pass for i in range(1, seq_len): # Broadcast score for all possible next tags broadcast_score = score.unsqueeze(2) # (batch_size, num_labels, 1) broadcast_emissions = emissions[i].unsqueeze(1) # (batch_size, 1, num_labels) # Compute next score next_score = broadcast_score + self.transitions + broadcast_emissions # Max and argmax next_score, indices = next_score.max(dim=1) # Save history and update score history.append(indices) score = torch.where(mask[i].unsqueeze(1), next_score, score) # Add end transitions score += self.end_transitions # Backtrack best_tags_list = [] for batch_idx in range(batch_size): # Find best last tag _, best_last_tag = score[batch_idx].max(dim=0) best_tags = [best_last_tag.item()] # Backtrack for hist in reversed(history): best_last_tag = hist[batch_idx][best_tags[-1]] best_tags.append(best_last_tag.item()) best_tags.reverse() best_tags_list.append(best_tags) return best_tags_list def compute_loss(self, emissions, tags, mask=None): """ Compute CRF negative log-likelihood loss Args: emissions: (batch_size, seq_len, num_labels) tags: (batch_size, seq_len) mask: (batch_size, seq_len) Returns: Negative log-likelihood loss """ if self.batch_first: emissions = emissions.transpose(0, 1) tags = tags.transpose(0, 1) if mask is not None: mask = mask.transpose(0, 1) seq_len, batch_size = tags.shape if mask is None: mask = torch.ones(seq_len, batch_size, dtype=torch.bool, device=emissions.device) # Compute score of gold sequence gold_score = self._compute_score(emissions, tags, mask) # Compute partition function (log sum of all possible sequences) forward_score = self._forward_algorithm(emissions, mask) # Loss = log(Z) - score(gold) loss = forward_score - gold_score return loss.mean() def _compute_score(self, emissions, tags, mask): """Compute score for a given tag sequence""" seq_len, batch_size = tags.shape score = self.start_transitions[tags[0]] score += emissions[0, torch.arange(batch_size), tags[0]] for i in range(1, seq_len): score += self.transitions[tags[i], tags[i-1]] * mask[i] score += emissions[i, torch.arange(batch_size), tags[i]] * mask[i] # Add end transitions last_tag_indices = mask.long().sum(0) - 1 last_tags = tags[last_tag_indices, torch.arange(batch_size)] score += self.end_transitions[last_tags] return score def _forward_algorithm(self, emissions, mask): """Forward algorithm to compute partition function""" seq_len, batch_size, num_labels = emissions.shape score = self.start_transitions + emissions[0] for i in range(1, seq_len): broadcast_score = score.unsqueeze(2) broadcast_emissions = emissions[i].unsqueeze(1) next_score = broadcast_score + self.transitions + broadcast_emissions next_score = torch.logsumexp(next_score, dim=1) score = torch.where(mask[i].unsqueeze(1), next_score, score) score += self.end_transitions return torch.logsumexp(score, dim=1)