mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
make ctx encoder vlm compatible
This commit is contained in:
parent
6f857a2dca
commit
2e88a18bd2
1 changed files with 32 additions and 3 deletions
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue