From fa3a97aa3612c5c759b70dc4165dd2c3907e815a Mon Sep 17 00:00:00 2001 From: Dan Yeh Date: Thu, 30 Jul 2026 22:07:53 +0800 Subject: [PATCH] export state_kv_projection for fusion append_kv in uzu --- lalamo/model_import/loaders/dflash_loader.py | 3 +++ lalamo/modules/speculators/dflash.py | 24 ++++++++++++++++++++ 2 files changed, 27 insertions(+) diff --git a/lalamo/model_import/loaders/dflash_loader.py b/lalamo/model_import/loaders/dflash_loader.py index 815996ea..029be45e 100644 --- a/lalamo/model_import/loaders/dflash_loader.py +++ b/lalamo/model_import/loaders/dflash_loader.py @@ -92,12 +92,14 @@ 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, ), @@ -105,6 +107,7 @@ def load_dflash_draft_model( ( context_projection, context_norm, + state_kv_projection, layers, output_norm, ), diff --git a/lalamo/modules/speculators/dflash.py b/lalamo/modules/speculators/dflash.py index 26ad517f..bd079ea0 100644 --- a/lalamo/modules/speculators/dflash.py +++ b/lalamo/modules/speculators/dflash.py @@ -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 ), @@ -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"],