Wire YARN rope_parameters through LFM2 and LFM2-MoE attention (#26187)
This commit is contained in:
@@ -125,13 +125,13 @@ class Lfm2Attention(nn.Module):
|
|||||||
if rope_parameters is not None and "rope_theta" in rope_parameters:
|
if rope_parameters is not None and "rope_theta" in rope_parameters:
|
||||||
rope_theta = rope_parameters["rope_theta"]
|
rope_theta = rope_parameters["rope_theta"]
|
||||||
else:
|
else:
|
||||||
rope_theta = config.rope_parameters["rope_theta"]
|
rope_theta = getattr(config, "rope_theta", 1000000.0)
|
||||||
|
|
||||||
self.rotary_emb = get_rope(
|
self.rotary_emb = get_rope(
|
||||||
head_size=self.head_dim,
|
head_size=self.head_dim,
|
||||||
rotary_dim=self.head_dim,
|
rotary_dim=self.head_dim,
|
||||||
max_position=getattr(config, "max_position_embeddings", 8192),
|
max_position=getattr(config, "max_position_embeddings", 8192),
|
||||||
rope_scaling=config.rope_parameters,
|
rope_scaling=rope_parameters or getattr(config, "rope_scaling", None),
|
||||||
base=rope_theta,
|
base=rope_theta,
|
||||||
is_neox_style=True,
|
is_neox_style=True,
|
||||||
dtype=torch.get_default_dtype(),
|
dtype=torch.get_default_dtype(),
|
||||||
|
|||||||
@@ -196,7 +196,7 @@ class Lfm2MoeAttention(nn.Module):
|
|||||||
head_size=self.head_dim,
|
head_size=self.head_dim,
|
||||||
rotary_dim=self.head_dim,
|
rotary_dim=self.head_dim,
|
||||||
max_position=getattr(config, "max_position_embeddings", 128000),
|
max_position=getattr(config, "max_position_embeddings", 128000),
|
||||||
rope_scaling=getattr(config, "rope_scaling", None),
|
rope_scaling=rope_parameters or getattr(config, "rope_scaling", None),
|
||||||
base=rope_theta,
|
base=rope_theta,
|
||||||
is_neox_style=True,
|
is_neox_style=True,
|
||||||
dtype=torch.get_default_dtype(),
|
dtype=torch.get_default_dtype(),
|
||||||
|
|||||||
Reference in New Issue
Block a user