""" Q-Former for GeneLingua Query-Former module that bridges DNA sequence embeddings with language understanding. Inspired by BLIP-2's Q-Former architecture. The Q-Former uses learnable query tokens to extract fixed-size, language-aligned features from variable-length DNA sequence embeddings (from DNABERT). """ import logging import math from typing import Optional, Tuple, List, Dict, Any import numpy as np logger = logging.getLogger(__name__) # Check for PyTorch try: import torch import torch.nn as nn import torch.nn.functional as F TORCH_AVAILABLE = True except ImportError: TORCH_AVAILABLE = False nn = None if TORCH_AVAILABLE: class MultiHeadAttention(nn.Module): """Multi-head attention layer.""" def __init__(self, embed_dim: int, num_heads: int, dropout: float = 0.1): super().__init__() self.embed_dim = embed_dim self.num_heads = num_heads self.head_dim = embed_dim // num_heads assert self.head_dim * num_heads == embed_dim, "embed_dim must be divisible by num_heads" self.q_proj = nn.Linear(embed_dim, embed_dim) self.k_proj = nn.Linear(embed_dim, embed_dim) self.v_proj = nn.Linear(embed_dim, embed_dim) self.out_proj = nn.Linear(embed_dim, embed_dim) self.dropout = nn.Dropout(dropout) self.scale = math.sqrt(self.head_dim) def forward( self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, attention_mask: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, torch.Tensor]: batch_size = query.size(0) # Project Q, K, V q = self.q_proj(query).view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2) k = self.k_proj(key).view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2) v = self.v_proj(value).view(batch_size, -1, self.num_heads, self.head_dim).transpose(1, 2) # Attention scores attn_weights = torch.matmul(q, k.transpose(-2, -1)) / self.scale if attention_mask is not None: attn_weights = attn_weights.masked_fill(attention_mask == 0, float('-inf')) attn_weights = F.softmax(attn_weights, dim=-1) attn_weights = self.dropout(attn_weights) # Apply attention attn_output = torch.matmul(attn_weights, v) attn_output = attn_output.transpose(1, 2).contiguous().view(batch_size, -1, self.embed_dim) attn_output = self.out_proj(attn_output) return attn_output, attn_weights class QFormerLayer(nn.Module): """Single Q-Former layer with self-attention and cross-attention.""" def __init__( self, embed_dim: int, num_heads: int, ff_dim: int, dropout: float = 0.1, ): super().__init__() # Self-attention for query tokens self.self_attn = MultiHeadAttention(embed_dim, num_heads, dropout) self.self_attn_norm = nn.LayerNorm(embed_dim) # Cross-attention to sequence embeddings self.cross_attn = MultiHeadAttention(embed_dim, num_heads, dropout) self.cross_attn_norm = nn.LayerNorm(embed_dim) # Feed-forward network self.ff = nn.Sequential( nn.Linear(embed_dim, ff_dim), nn.GELU(), nn.Dropout(dropout), nn.Linear(ff_dim, embed_dim), nn.Dropout(dropout), ) self.ff_norm = nn.LayerNorm(embed_dim) def forward( self, query_tokens: torch.Tensor, sequence_embeddings: torch.Tensor, sequence_mask: Optional[torch.Tensor] = None, ) -> torch.Tensor: # Self-attention residual = query_tokens query_tokens = self.self_attn_norm(query_tokens) query_tokens, _ = self.self_attn(query_tokens, query_tokens, query_tokens) query_tokens = residual + query_tokens # Cross-attention to sequence residual = query_tokens query_tokens = self.cross_attn_norm(query_tokens) query_tokens, _ = self.cross_attn(query_tokens, sequence_embeddings, sequence_embeddings, sequence_mask) query_tokens = residual + query_tokens # Feed-forward residual = query_tokens query_tokens = self.ff_norm(query_tokens) query_tokens = self.ff(query_tokens) query_tokens = residual + query_tokens return query_tokens class QFormer(nn.Module): """ Query-Former that bridges DNA embeddings with language understanding. Architecture: - Learnable query tokens (num_queries x embed_dim) - Multiple Q-Former layers with self and cross attention - Projects DNA sequence embeddings to language-aligned space Usage: qformer = QFormer(num_queries=32, embed_dim=768) dna_embeddings = dnabert.embed_batch(sequences) # (batch, seq_len, 768) language_features = qformer(dna_embeddings) # (batch, num_queries, embed_dim) """ def __init__( self, num_queries: int = 32, embed_dim: int = 768, num_layers: int = 6, num_heads: int = 12, ff_dim: int = 3072, dropout: float = 0.1, dna_embed_dim: int = 768, ): """ Initialize Q-Former. Args: num_queries: Number of learnable query tokens embed_dim: Embedding dimension num_layers: Number of Q-Former layers num_heads: Number of attention heads ff_dim: Feed-forward hidden dimension dropout: Dropout rate dna_embed_dim: DNABERT embedding dimension """ super().__init__() self.num_queries = num_queries self.embed_dim = embed_dim # Learnable query tokens self.query_tokens = nn.Parameter(torch.randn(1, num_queries, embed_dim) * 0.02) # Project DNA embeddings to Q-Former dimension if needed self.dna_projection = None if dna_embed_dim != embed_dim: self.dna_projection = nn.Linear(dna_embed_dim, embed_dim) # Q-Former layers self.layers = nn.ModuleList([ QFormerLayer(embed_dim, num_heads, ff_dim, dropout) for _ in range(num_layers) ]) self.final_norm = nn.LayerNorm(embed_dim) # Output projection (optional, for specific downstream tasks) self.output_projection = nn.Linear(embed_dim, embed_dim) self.logger = logging.getLogger("QFormer") def forward( self, sequence_embeddings: torch.Tensor, sequence_mask: Optional[torch.Tensor] = None, return_attention: bool = False, ) -> torch.Tensor: """ Extract language-aligned features from DNA sequence embeddings. Args: sequence_embeddings: DNA embeddings from DNABERT (batch, seq_len, dna_dim) sequence_mask: Attention mask for sequence (batch, seq_len) return_attention: Whether to return attention weights Returns: Query features (batch, num_queries, embed_dim) """ batch_size = sequence_embeddings.size(0) # Project DNA embeddings if needed if self.dna_projection is not None: sequence_embeddings = self.dna_projection(sequence_embeddings) # Expand query tokens for batch query_tokens = self.query_tokens.expand(batch_size, -1, -1) # Process through Q-Former layers for layer in self.layers: query_tokens = layer(query_tokens, sequence_embeddings, sequence_mask) # Final normalization query_tokens = self.final_norm(query_tokens) return query_tokens def get_pooled_output(self, sequence_embeddings: torch.Tensor) -> torch.Tensor: """ Get single pooled representation from sequence. Args: sequence_embeddings: DNA embeddings (batch, seq_len, dna_dim) Returns: Pooled features (batch, embed_dim) """ query_features = self.forward(sequence_embeddings) # Mean pool over query tokens pooled = query_features.mean(dim=1) return self.output_projection(pooled) class DNALanguageBridge(nn.Module): """ Complete bridge between DNA sequences and language models. Combines: - DNABERT for sequence encoding - Q-Former for feature extraction - Projection to language model space """ def __init__( self, dna_embedder=None, num_queries: int = 32, qformer_layers: int = 6, language_dim: int = 768, freeze_dna_encoder: bool = True, ): """ Initialize DNA-Language bridge. Args: dna_embedder: DNABERT embedder instance num_queries: Number of Q-Former queries qformer_layers: Number of Q-Former layers language_dim: Language model dimension freeze_dna_encoder: Whether to freeze DNABERT """ super().__init__() self.dna_embedder = dna_embedder self.freeze_dna_encoder = freeze_dna_encoder # Get DNABERT embedding dimension if dna_embedder is not None: dna_dim = dna_embedder.get_embedding_dim() else: dna_dim = 768 # Q-Former self.qformer = QFormer( num_queries=num_queries, embed_dim=language_dim, num_layers=qformer_layers, dna_embed_dim=dna_dim, ) # Language projection self.language_projection = nn.Linear(language_dim, language_dim) self.logger = logging.getLogger("DNALanguageBridge") def encode_sequence(self, sequence: str) -> torch.Tensor: """ Encode a DNA sequence to language-aligned features. Args: sequence: DNA sequence string Returns: Language-aligned features (1, num_queries, language_dim) """ return self.encode_sequences([sequence]) def encode_sequences(self, sequences: List[str]) -> torch.Tensor: """ Encode multiple DNA sequences. Args: sequences: List of DNA sequences Returns: Language-aligned features (batch, num_queries, language_dim) """ # Get DNA embeddings if self.dna_embedder is not None: if hasattr(self.dna_embedder, '_load_model'): self.dna_embedder._load_model() # Get hidden states (not just pooled output) # This requires access to the transformer outputs embeddings = self._get_sequence_hidden_states(sequences) else: # Fallback: create dummy embeddings embeddings = torch.randn(len(sequences), 128, 768) # Q-Former processing if self.freeze_dna_encoder: with torch.no_grad(): query_features = self.qformer(embeddings) else: query_features = self.qformer(embeddings) return query_features def _get_sequence_hidden_states(self, sequences: List[str]) -> torch.Tensor: """Get full hidden states from DNABERT (not just mean pooled).""" embedder = self.dna_embedder embedder._load_model() import torch as th # Preprocess processed = [embedder._preprocess_sequence(seq) for seq in sequences] # Tokenize inputs = embedder.tokenizer( processed, return_tensors="pt", padding=True, truncation=True, max_length=embedder.max_length, ) inputs = {k: v.to(embedder.device) for k, v in inputs.items()} # Get hidden states with th.no_grad(): outputs = embedder.model(**inputs, output_hidden_states=True) return outputs.last_hidden_state def get_text_compatible_embeddings(self, sequences: List[str]) -> torch.Tensor: """ Get embeddings suitable for text-based retrieval. Args: sequences: DNA sequences Returns: Pooled embeddings (batch, language_dim) """ query_features = self.encode_sequences(sequences) pooled = query_features.mean(dim=1) return self.language_projection(pooled) class QFormerConfig: """Configuration for Q-Former.""" def __init__( self, num_queries: int = 32, embed_dim: int = 768, num_layers: int = 6, num_heads: int = 12, ff_dim: int = 3072, dropout: float = 0.1, dna_embed_dim: int = 768, language_dim: int = 768, ): self.num_queries = num_queries self.embed_dim = embed_dim self.num_layers = num_layers self.num_heads = num_heads self.ff_dim = ff_dim self.dropout = dropout self.dna_embed_dim = dna_embed_dim self.language_dim = language_dim def to_dict(self) -> Dict[str, Any]: return vars(self) @classmethod def from_dict(cls, config_dict: Dict[str, Any]) -> "QFormerConfig": return cls(**config_dict) # Fallback for when PyTorch is not available class FallbackQFormer: """Simple fallback Q-Former using numpy when PyTorch is not available.""" def __init__(self, num_queries: int = 32, embed_dim: int = 768): self.num_queries = num_queries self.embed_dim = embed_dim self.query_tokens = np.random.randn(num_queries, embed_dim) * 0.02 self.logger = logging.getLogger("FallbackQFormer") self.logger.warning("Using fallback Q-Former (PyTorch not available)") def forward(self, sequence_embeddings: np.ndarray) -> np.ndarray: """ Simple attention-based feature extraction. Args: sequence_embeddings: (batch, seq_len, embed_dim) Returns: Query features (batch, num_queries, embed_dim) """ batch_size = sequence_embeddings.shape[0] results = [] for i in range(batch_size): seq_emb = sequence_embeddings[i] # (seq_len, embed_dim) # Simple attention: dot product between queries and sequence attn_scores = np.dot(self.query_tokens, seq_emb.T) # (num_queries, seq_len) attn_weights = self._softmax(attn_scores, axis=-1) # Weighted sum output = np.dot(attn_weights, seq_emb) # (num_queries, embed_dim) results.append(output) return np.stack(results) def _softmax(self, x: np.ndarray, axis: int = -1) -> np.ndarray: exp_x = np.exp(x - np.max(x, axis=axis, keepdims=True)) return exp_x / np.sum(exp_x, axis=axis, keepdims=True) def get_pooled_output(self, sequence_embeddings: np.ndarray) -> np.ndarray: """Get pooled representation.""" query_features = self.forward(sequence_embeddings) return np.mean(query_features, axis=1) def get_qformer(config: QFormerConfig = None, use_pytorch: bool = True): """ Get Q-Former instance. Args: config: Q-Former configuration use_pytorch: Try to use PyTorch version Returns: Q-Former instance """ if config is None: config = QFormerConfig() if use_pytorch and TORCH_AVAILABLE: return QFormer( num_queries=config.num_queries, embed_dim=config.embed_dim, num_layers=config.num_layers, num_heads=config.num_heads, ff_dim=config.ff_dim, dropout=config.dropout, dna_embed_dim=config.dna_embed_dim, ) return FallbackQFormer( num_queries=config.num_queries, embed_dim=config.embed_dim, ) if __name__ == "__main__": logging.basicConfig(level=logging.INFO) print("Testing Q-Former") print("=" * 50) config = QFormerConfig(num_queries=8, num_layers=2) qformer = get_qformer(config) print(f"Q-Former type: {type(qformer).__name__}") print(f"Num queries: {config.num_queries}") print(f"Embed dim: {config.embed_dim}") # Test with dummy data if TORCH_AVAILABLE: batch_size = 2 seq_len = 64 dummy_embeddings = torch.randn(batch_size, seq_len, config.dna_embed_dim) output = qformer(dummy_embeddings) print(f"\nInput shape: {dummy_embeddings.shape}") print(f"Output shape: {output.shape}") print(f"Expected: ({batch_size}, {config.num_queries}, {config.embed_dim})") else: batch_size = 2 seq_len = 64 dummy_embeddings = np.random.randn(batch_size, seq_len, config.dna_embed_dim) output = qformer.forward(dummy_embeddings) print(f"\nInput shape: {dummy_embeddings.shape}") print(f"Output shape: {output.shape}") print("\nQ-Former test complete!")