diff --git a/i6_models/parts/conformer/mhsa_rel_pos.py b/i6_models/parts/conformer/mhsa_rel_pos.py index 6f05dc26..821b1c10 100644 --- a/i6_models/parts/conformer/mhsa_rel_pos.py +++ b/i6_models/parts/conformer/mhsa_rel_pos.py @@ -163,15 +163,13 @@ def forward(self, input_tensor: torch.Tensor, sequence_mask: torch.Tensor) -> to k = key_seq.view(batch_dim_size, -1, self.num_heads, self.embed_dim_per_head) # [B, T', #heads, F'] if self.learnable_pos_emb: - pos_seq_q = torch.arange(time_dim_size, device=input_tensor.device) - pos_seq_k = torch.arange(time_dim_size, device=input_tensor.device) - - distance_mat = pos_seq_k[None, :] - pos_seq_q[:, None] - distance_mat_clipped = torch.clamp(distance_mat, -self.rel_pos_clip, self.rel_pos_clip) - - final_mat = distance_mat_clipped + self.rel_pos_clip - - rel_pos_embeddings = self.rel_pos_embeddings[final_mat] # [T, T', pos_emb_dim] + # 1D optimization: 2T-1 unique relative positions instead of T×T distance matrix. + # Build [-(T-1), ..., -1, 0, 1, ..., T-1] directly. + rel_pos = torch.arange(-(time_dim_size - 1), time_dim_size, device=input_tensor.device) + indices = torch.clamp(rel_pos, -self.rel_pos_clip, self.rel_pos_clip) + self.rel_pos_clip + rel_pos_embeddings = self.rel_pos_embeddings[indices].view( + 1, 2 * time_dim_size - 1, self.pos_emb_dim + ) # [1, T+T'-1, pos_emb_dim] else: rel_pos_embeddings = ( self._sinusoidal_pe( @@ -207,9 +205,7 @@ def forward(self, input_tensor: torch.Tensor, sequence_mask: torch.Tensor) -> to q_with_bias_v, rel_pos_embeddings.to(device=q_with_bias_v.device, dtype=q_with_bias_v.dtype), ) # [B, #heads, T, T'] or [B, #heads, T, T+T'+1] - if not self.learnable_pos_emb: - attn_bd = self._rel_shift_bhij(attn_bd, k_len=time_dim_size) # [B, #heads, T, T'] - + attn_bd = self._rel_shift_bhij(attn_bd, k_len=time_dim_size) # [B, #heads, T, T'] # We use attn_mask to add BD matrix to attention scores. # # Inside torch's SDPA the mask is added after regular scaling, so to get correct