From 0e05644d6a66a986c42c79f2652e372063b4bf35 Mon Sep 17 00:00:00 2001 From: 51616 Date: Fri, 12 Sep 2025 19:05:05 +0900 Subject: [PATCH] iterative mode for hypernet --- run_eval.py | 5 + src/ctx_to_lora/eval_utils.py | 3 + src/ctx_to_lora/modeling/aggregator.py | 129 +++---------------------- src/ctx_to_lora/modeling/hypernet.py | 38 +++++--- 4 files changed, 49 insertions(+), 126 deletions(-) diff --git a/run_eval.py b/run_eval.py index 742ec12..0ce8c62 100644 --- a/run_eval.py +++ b/run_eval.py @@ -99,6 +99,11 @@ if __name__ == "__main__": default=None, help="Number of generated queries to use for context distillation.", ) + parser.add_argument( + "--use_iterative_mode", + action="store_true", + help="Use iterative mode LoRA layer-by-layer generation", + ) cli_args = vars(parser.parse_args()) # setup_logging(output_dir, debug=os.getenv("DEBUG", False)) diff --git a/src/ctx_to_lora/eval_utils.py b/src/ctx_to_lora/eval_utils.py index 2484909..0edf9cb 100644 --- a/src/ctx_to_lora/eval_utils.py +++ b/src/ctx_to_lora/eval_utils.py @@ -771,6 +771,7 @@ def evaluate( use_flash_attn=True, use_sequence_packing=False, # for generation ) + model.enable_iterative_mode(args.use_iterative_mode) add_tracker(model.base_model.generate, "generate") add_tracker(model.generate_weights, "generate_weights") add_tracker(model.combine_lora, "combine_lora") @@ -1004,6 +1005,7 @@ def run_eval( cd_update_iterations: int = 10, cd_use_gen_q: bool = False, num_gen_q: int = 20, + use_iterative_mode: bool = False, ) -> None: """Run evaluation with the specified parameters.""" assert bool(model_name_or_path) ^ bool(checkpoint_path), ( @@ -1047,6 +1049,7 @@ def run_eval( # modulated model doesn't see ctx by default # but remove_context has to be false for correct file naming args.remove_context = False + args.use_iterative_mode = use_iterative_mode else: args = Namespace( model_name_or_path=model_name_or_path, diff --git a/src/ctx_to_lora/modeling/aggregator.py b/src/ctx_to_lora/modeling/aggregator.py index 3a472bb..057f69b 100644 --- a/src/ctx_to_lora/modeling/aggregator.py +++ b/src/ctx_to_lora/modeling/aggregator.py @@ -98,6 +98,7 @@ class Perceiver(nn.Module): **kwargs, ): super().__init__() + assert num_extra_modules == 0 self.num_layers = num_layers self.num_modules = num_modules self.num_extra_modules = num_extra_modules @@ -133,6 +134,10 @@ class Perceiver(nn.Module): attn_implementation="flash_attention_2", ) self.perceiver = Idefics2Perceiver(self.config, self.decoder_config) + self.iterative_mode = False + + def enable_iterative_mode(self, x: bool): + self.iterative_mode = x def forward( self, @@ -141,30 +146,7 @@ class Perceiver(nn.Module): ctx_attn_mask: Integer[Tensor, "bs seq_len"] | None = None, ctx_position_ids: Integer[Tensor, "bs seq_len"] | None = None, ): - # if len(ctx_features.shape) == 3: - # # ["bs seq_len feature_dim"] to - # # ["bs n_output_queries feature_dim"] - # x = self.perceiver(ctx_features, ctx_attn_mask, ctx_position_ids) - # elif len(ctx_features.shape) == 4: - # # ["bs n_ctx_model_layers seq_len feature_dim"] to - # # to ["bs n_ctx_model_layers n_output_queries feature_dim"] - # x = self.perceiver(ctx_features, ctx_attn_mask, ctx_position_ids) - # x = rearrange( - # x, - # ( - # "bs n_ctx_model_layers (n_modules r) seq_len d -> " - # "bs (n_ctx_model_layers n_modules r) d" - # ), - # n_ctx_model_layers=self.n_ctx_model_layers, - # n_modules=self.num_modules, - # r=self.r, - # ) - # else: - # raise ValueError( - # f"ctx_features should be 3D or 4D tensor, got {ctx_features.shape}" - # ) - - if self.layer_to_layer: + if self.layer_to_layer and not self.iterative_mode: if ctx_attn_mask is not None: ctx_attn_mask = repeat( ctx_attn_mask, @@ -186,15 +168,17 @@ class Perceiver(nn.Module): "1 num_layers seq_len feature_dim -> 1 (num_layers seq_len) feature_dim", ) - # print(f"ctx_features shape: {ctx_features.shape}") - # print( - # f"ctx_attn_mask shape: {ctx_attn_mask.shape if ctx_attn_mask is not None else None}" - # ) - # print( - # f"ctx_position_ids shape: {ctx_position_ids.shape if ctx_position_ids is not None else None}" - # ) x = self.perceiver(ctx_features, ctx_attn_mask, ctx_position_ids) + if self.layer_to_layer and self.iterative_mode: + lora_x = rearrange( + x, + "bs (n_modules r) d -> bs n_modules r d", + n_modules=self.num_modules, + r=self.r, + ) + return lora_x, None + if self.layer_to_layer: per_layer_size = self.num_modules * self.r + self.num_extra_modules x = rearrange( @@ -228,93 +212,10 @@ class Perceiver(nn.Module): n_layers=self.num_layers, ) - # x = rearrange( - # x, - # "bs (n_layers n_modules) d -> bs n_layers n_modules d", - # n_modules=self.num_modules, - # n_layers=self.num_layers, - # ) - # lora_emb, extra_emb = unpack( - # emb, - # [[self.num_modules], [self.num_extra_modules]], - # "bs n_layers * feature_dim", - # ) return lora_x, extra_x -# class Pooler(nn.Module): -# def __init__( -# self, -# feature_size: int, -# output_size: int, -# pooling_type: POOL_FN, -# num_layers: int, -# num_modules: int, -# *args, -# **kwargs, -# ): -# super().__init__() -# self.num_layers = num_layers -# self.num_modules = num_modules - -# # NOTE: features will be projected to size = output_size // 2 -# # then cat with layer and module embeddings (each with size output_size // 4) -# # which are collectively form features with size = output_size -# self.pool_fn = get_pooling_fn(pooling_type) -# self.feature_proj = nn.Linear(feature_size, output_size // 2) -# self.ln = nn.LayerNorm(output_size // 2) -# self.layer_embs = nn.Sequential( -# nn.Embedding(num_layers, output_size // 4), -# nn.LayerNorm(output_size // 4), -# ) -# self.module_embs = nn.Sequential( -# nn.Embedding(num_modules, output_size // 4), -# nn.LayerNorm(output_size // 4), -# ) -# self.mixer = Mixer(output_size, output_size * 4, output_size) -# self.mlp = MLPResidualBlock(output_size, output_size * 4, output_size) - -# self.register_buffer("layer_indices", torch.arange(num_layers)) -# self.register_buffer("module_indices", torch.arange(num_modules)) - -# def forward( -# self, -# features: Float[Tensor, "bs seq_len feature_dim"], -# attn_mask: Integer[Tensor, "bs seq_len"] | None = None, -# ): -# bs = features.shape[0] - -# # [bs, feature_dim] -# x = self.ln(self.feature_proj(self.pool_fn(features, attn_mask).float())) -# x = repeat( -# x, -# "bs d -> bs n_layers n_modules d", -# n_modules=self.num_modules, -# n_layers=self.num_layers, -# ) - -# layer_embs = self.layer_embs(self.layer_indices) # [num_layers, d] -# layer_embs = repeat( -# layer_embs, -# "n_layers d -> bs n_layers n_modules d", -# bs=bs, -# n_modules=self.num_modules, -# ) - -# module_embs = self.module_embs(self.module_indices) # [num_modules, d] -# module_embs = repeat( -# module_embs, -# "n_modules d -> bs n_layers n_modules d", -# bs=bs, -# n_layers=self.num_layers, -# ) - -# emb = torch.cat([x, layer_embs, module_embs], dim=3) -# return self.mlp(self.mixer(emb)) - - AGGREGATOR_CLS = { - # AGGREGATOR_TYPE.POOLER: Pooler, AGGREGATOR_TYPE.PERCEIVER: Perceiver, } diff --git a/src/ctx_to_lora/modeling/hypernet.py b/src/ctx_to_lora/modeling/hypernet.py index bac740b..4f61e7c 100644 --- a/src/ctx_to_lora/modeling/hypernet.py +++ b/src/ctx_to_lora/modeling/hypernet.py @@ -644,6 +644,10 @@ class HyperLoRA(nn.Module): ) return {k: v for k, v in zip(self.extra_modules, layernorms)} + def enable_iterative_mode(self, x: bool): + self.iterative_mode = x + self.aggregator.enable_iterative_mode(x) + def forward( self, features: Float[Tensor, "bs seq_len feature_dim"], @@ -653,15 +657,22 @@ class HyperLoRA(nn.Module): ): # [bs, n_layers, n_total_modules, r, feature_dim] with torch.autocast(device_type="cuda", dtype=torch.bfloat16): - lora_emb, extra_emb = self.aggregator(features, attn_mask, position_ids) + if self.aggregator.layer_to_layer and self.iterative_mode: + # iterative inference + # features: [bs num_layers seq_len feature_dim] + bs, n_layers = features.shape[0:2] + lora_emb = torch.empty( + (bs, n_layers, self.num_modules, self.r, self.config.latent_size), + device=features.device, + ) + for i in range(n_layers): + lora_emb[:, i], _ = self.aggregator( + features[:, i], attn_mask, position_ids + ) - # TODO: add pos emb - # TODO: - # concat queries from the same chunks - # do self-attn - # unpack back to [bs, n_layers, n_total_modules, r, feature_dim] - # here bs = sum(n_chunks) - # we then + else: + # batched inference + lora_emb, _ = self.aggregator(features, attn_mask, position_ids) # [bs, n_layers, n_modules, r, max_in_d_outim] flat_loras = None @@ -672,10 +683,10 @@ class HyperLoRA(nn.Module): flat_loras = self.head(norm_lora_emb) flat_layernorms = None - if self.num_extra_modules: - # [bs, n_layers, n_extra_modules, base_hidden_size] - extra_emb = self.extra_layers(extra_emb) - flat_layernorms = self.extra_head(extra_emb) + # if self.num_extra_modules: + # # [bs, n_layers, n_extra_modules, base_hidden_size] + # extra_emb = self.extra_layers(extra_emb) + # flat_layernorms = self.extra_head(extra_emb) return flat_loras, flat_layernorms @@ -952,6 +963,9 @@ class ModulatedPretrainedModel(nn.Module): ctx_features, ctx_attn_mask, ctx_position_ids ) + def enable_iterative_mode(self, x: bool): + self.hypernet.enable_iterative_mode(x) + def forward( self, ctx_ids: Integer[Tensor, "n_ctx ctx_len"] | None = None,