fix(lumina): pass to_compute_mask, not use_cache, to LLaDABlock.attention

attention() takes (q, k, v, attention_bias, layer_past, to_compute_mask) and
has no use_cache parameter. Three of the four call sites still pass
use_cache=, which raises TypeError; only LLaDALlamaBlock's non-checkpointed
branch -- the path the shipped block_type=llama config takes -- is correct.
This commit is contained in:
Anai-Guo
2026-08-30 18:20:30 -07:00
parent 9f8e45c69b
commit d556544b30
+3 -3
View File
@@ -874,10 +874,10 @@ class LLaDASequentialBlock(LLaDABlock):
if self._activation_checkpoint_fn is not None:
att, cache = self._activation_checkpoint_fn( # type: ignore
self.attention, q, k, v, attention_bias, layer_past=layer_past, use_cache=use_cache
self.attention, q, k, v, attention_bias, layer_past=layer_past
)
else:
att, cache = self.attention(q, k, v, attention_bias, layer_past=layer_past, use_cache=use_cache)
att, cache = self.attention(q, k, v, attention_bias, layer_past=layer_past)
x = x + self.dropout(att)
@@ -967,7 +967,7 @@ class LLaDALlamaBlock(LLaDABlock):
if self._activation_checkpoint_fn is not None:
att, cache = self._activation_checkpoint_fn( # type: ignore
self.attention, q, k, v, attention_bias, layer_past=layer_past, use_cache=use_cache
self.attention, q, k, v, attention_bias, layer_past=layer_past, to_compute_mask=to_compute_mask
)
else:
att, cache = self.attention(q, k, v, attention_bias, layer_past=layer_past, to_compute_mask=to_compute_mask)