From 2e88a18bd21e897e06057bd66ddf6822edac6dd2 Mon Sep 17 00:00:00 2001 From: 51616 Date: Mon, 10 Nov 2025 16:54:55 +0900 Subject: [PATCH] make ctx encoder vlm compatible --- src/ctx_to_lora/modeling/ctx_encoder.py | 35 ++++++++++++++++++++++--- 1 file changed, 32 insertions(+), 3 deletions(-) diff --git a/src/ctx_to_lora/modeling/ctx_encoder.py b/src/ctx_to_lora/modeling/ctx_encoder.py index 682edee..532bed6 100644 --- a/src/ctx_to_lora/modeling/ctx_encoder.py +++ b/src/ctx_to_lora/modeling/ctx_encoder.py @@ -91,7 +91,11 @@ class EmbeddingOnly(nn.Module): class PerLayerActivations(nn.Module): def __init__(self, base_model: PreTrainedModel, config: CtxEncoderArguments): super().__init__() - base_model = get_base_model(base_model) # remove lm head + self.keep_lm_head = getattr(config, "keep_lm_head", False) + if not self.keep_lm_head: + base_model = get_base_model(base_model) # remove lm head + else: + base_model.lm_head = nn.Identity() # -1 to remove last attn block if config.ctx_encoder_last_layer is not None: @@ -99,13 +103,34 @@ class PerLayerActivations(nn.Module): else: last_layer = -1 - base_model.layers = base_model.layers[:last_layer] + if self.keep_lm_head: + base_model.model.layers = base_model.model.layers[:last_layer] + else: + base_model.layers = base_model.layers[:last_layer] self.base_model = base_model @property def config(self): return self.base_model.config + def get_input_embeddings(self): + return self.base_model.get_input_embeddings() + + def set_input_embeddings(self, value): + self.base_model.set_input_embeddings(value) + + def get_output_embeddings(self): + return self.base_model.get_output_embeddings() + + def set_output_embeddings(self, new_embeddings): + self.base_model.set_output_embeddings(new_embeddings) + + def set_decoder(self, decoder): + self.base_model.set_decoder(decoder) + + def get_decoder(self): + return self.base_model.get_decoder() + @torch.no_grad() def forward(self, **kwargs): kwargs["output_hidden_states"] = True # Force output of hidden states @@ -113,7 +138,11 @@ class PerLayerActivations(nn.Module): # Return all layers' activations except the last one # from embeddings to the input of the last attn block # Shape: (batch_size, num_layers, seq_len, hidden_size) - return torch.stack(outputs.hidden_states, dim=1) + + if self.keep_lm_head: + return outputs + else: + return torch.stack(outputs.hidden_states, dim=1) class CTX_ENCODER_TYPE(str, Enum):