Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions lalamo/model_import/loaders/dflash_loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -92,19 +92,22 @@ def load_dflash_draft_model(
)
for layer_index, layer in enumerate(module.layers)
)
state_kv_projection = module.state_kv_projection_from_layers(layers)
output_norm = load_rmsnorm(module.output_norm, weights_dict, path / "norm")

return load_as_at(
lambda draft_model: (
draft_model.context_projection,
draft_model.context_norm,
draft_model.state_kv_projection,
draft_model.layers,
draft_model.output_norm,
),
module,
(
context_projection,
context_norm,
state_kv_projection,
layers,
output_norm,
),
Expand Down
24 changes: 24 additions & 0 deletions lalamo/modules/speculators/dflash.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,6 +142,15 @@ def init(self, initializer: Initializer) -> "DFlashDraftModel":
),
context_norm=self.context_norm_config.init(initializer, self.model_dim),
rope=self.rope_config.init(initializer),
state_kv_projection=LinearConfig().init(
initializer,
self.model_dim,
tuple(
2 * layer.attention_config.num_groups * layer.attention_config.head_dim
for layer in self.layer_configs
),
has_biases=False,
),
layers=tuple(
layer_config.init(initializer, self.model_dim, self.hidden_dim) for layer_config in self.layer_configs
),
Expand Down Expand Up @@ -211,9 +220,24 @@ class DFlashDraftModel(LalamoModule[DFlashDraftConfig]):
context_projection: Linear
context_norm: Normalization
rope: RoPE
state_kv_projection: Linear
layers: tuple[DFlashDraftLayer, ...]
output_norm: Normalization

def state_kv_projection_from_layers(self, layers: tuple[DFlashDraftLayer, ...]) -> Linear:
qkv_projections = tuple(layer.attention.qkv_projection for layer in layers)
key_value_weights = jnp.concatenate(
tuple(projection.weights.decompress()[projection.output_dims[0] :] for projection in qkv_projections),
axis=0,
)
weights = qkv_projections[0].weights.spec.compress(
key_value_weights,
key=jax.random.key(0),
sharding_config=self.state_kv_projection.weights.sharding_config,
is_sharded=self.state_kv_projection.weights.is_sharded,
)
return eqx.tree_at(lambda projection: projection.weights, self.state_kv_projection, weights)

def positional_embeddings(
self,
token_positions: Int[Array, "batch tokens"],
Expand Down
Loading