make ctx encoder vlm compatible

This commit is contained in:
51616 2025-11-10 16:54:55 +09:00
parent 6f857a2dca
commit 2e88a18bd2

View file

@ -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):