mrs83 commited on
Commit
ce41e57
·
verified ·
1 Parent(s): d8d787b

Upload folder using huggingface_hub

Browse files
Files changed (1) hide show
  1. modeling_echo.py +12 -6
modeling_echo.py CHANGED
@@ -462,11 +462,6 @@ class SlidingWindowAttention(nn.Module):
462
  self.head_dim = self.hidden_size // self.num_heads
463
  self.window_size = getattr(config, "window_size", 128)
464
  self.attention_masking = getattr(config, "attention_masking", "causal")
465
- # vLLM's Transformers backend inspects `module.is_causal` to pick the
466
- # decoder (causal) vs encoder (bidirectional) attention backend —
467
- # mirror the masking mode so non_causal_window models keep their
468
- # bidirectional semantics under vLLM.
469
- self.is_causal = self.attention_masking != "non_causal_window"
470
 
471
  self.qkv_proj = nn.Linear(self.hidden_size, 3 * self.hidden_size, bias=False)
472
  self.out_proj = nn.Linear(self.hidden_size, self.hidden_size, bias=False)
@@ -569,7 +564,17 @@ class SlidingWindowAttention(nn.Module):
569
  # Replace -inf with 0 for the permitted window (float mask expected by sdpa)
570
  mask = torch.where(mask == float("-inf"), mask, torch.zeros_like(mask))
571
 
572
- y = self._attn_call(attn_fn, q, k, v, mask.unsqueeze(0).unsqueeze(0), is_causal=False)
 
 
 
 
 
 
 
 
 
 
573
  else:
574
  # Decoding: Recurrent step, attend only to the last window_size tokens
575
  y = self._attn_call(attn_fn, q, k_attn, v_attn, None, is_causal=False)
@@ -1386,6 +1391,7 @@ class EchoModelForPooling(EchoModel):
1386
  past_key_values=past_key_values,
1387
  inputs_embeds=inputs_embeds,
1388
  position_ids=position_ids,
 
1389
  output_dsrn_telemetry=output_dsrn_telemetry,
1390
  output_attentions=output_attentions,
1391
  output_hidden_states=output_hidden_states,
 
462
  self.head_dim = self.hidden_size // self.num_heads
463
  self.window_size = getattr(config, "window_size", 128)
464
  self.attention_masking = getattr(config, "attention_masking", "causal")
 
 
 
 
 
465
 
466
  self.qkv_proj = nn.Linear(self.hidden_size, 3 * self.hidden_size, bias=False)
467
  self.out_proj = nn.Linear(self.hidden_size, self.hidden_size, bias=False)
 
564
  # Replace -inf with 0 for the permitted window (float mask expected by sdpa)
565
  mask = torch.where(mask == float("-inf"), mask, torch.zeros_like(mask))
566
 
567
+ # Block attention to padding tokens. The input attention_mask is
568
+ # only present on the plain transformers/ST path (padded batches);
569
+ # vLLM's pooling runner never passes one (its flattened batches
570
+ # are unpadded), so this is a no-op there.
571
+ mask4 = mask.unsqueeze(0).unsqueeze(0) # (1, 1, T, kv)
572
+ input_mask = kwargs.get("attention_mask")
573
+ if input_mask is not None and self.attention_masking == "non_causal_window":
574
+ key_blocked = input_mask.bool().unsqueeze(1).unsqueeze(1) # (B, 1, 1, T)
575
+ mask4 = mask4.masked_fill(~key_blocked, float("-inf"))
576
+
577
+ y = self._attn_call(attn_fn, q, k, v, mask4, is_causal=False)
578
  else:
579
  # Decoding: Recurrent step, attend only to the last window_size tokens
580
  y = self._attn_call(attn_fn, q, k_attn, v_attn, None, is_causal=False)
 
1391
  past_key_values=past_key_values,
1392
  inputs_embeds=inputs_embeds,
1393
  position_ids=position_ids,
1394
+ attention_mask=attention_mask,
1395
  output_dsrn_telemetry=output_dsrn_telemetry,
1396
  output_attentions=output_attentions,
1397
  output_hidden_states=output_hidden_states,