embedding encoder + mean agg

This commit is contained in:
51616 2024-12-20 15:09:32 +00:00
parent f2606dc29d
commit 8febadd096

View file

@ -1,75 +1,158 @@
import logging
from dataclasses import dataclass, field from dataclasses import dataclass, field
from enum import Enum
from functools import partial
from typing import Any, Optional, Tuple, Union from typing import Any, Optional, Tuple, Union
import torch import torch
from torch import nn from jaxtyping import Float, Integer
from peft import LoraConfig
from torch import Tensor, nn
from transformers import PreTrainedModel from transformers import PreTrainedModel
from transformers.modeling_outputs import ModelOutput from transformers.modeling_outputs import ModelOutput
# @dataclass from model_loading import get_lora_config, get_model_and_tokenizer
# class ModelArguments:
# model_name_or_path: str = field(default="mistralai/Mistral-7B-v0.1") logger = logging.getLogger(__name__)
# lora_r: int = field(default=128, metadata={"help": "lora rank"})
# lora_dropout: float = field(default=0.05, metadata={"help": "lora dropout"})
# train: bool = field(
# default=True,
# metadata={
# "help": "if true, the model ckpt will be initialized for training; else, it's for inference"
# },
# )
class ModulatedPretrainedModel(nn.Module): def inv_bool_mask(m: Integer[Tensor, "bs seq_len"]) -> Integer[Tensor, "bs seq_len 1"]:
# def __init__(self, model_args, training_args, lora_config): return (m - 1).bool().unsqueeze(-1)
# super().__init__()
# self.model_args = model_args
# self.training_args = training_args # TODO: implement Perceiver
# self.model_name = model_args.model_name_or_path class Perceiver(nn.Module): ...
# self.lora_config = lora_config
# self.model = AutoModelForCausalLM.from_pretrained(
# self.model_name, class MeanPool(nn.Module):
# torch_dtype=torch.float16 if training_args.bf16 is False else torch.bfloat16, def forward(
# use_flash_attention_2=True, self,
# resume_download=True, features: Float[Tensor, "bs seq_len feature_dim"],
# ) attn_mask: Optional[Integer[Tensor, "bs seq_len"]] = None,
def __init__(self, base_model: PreTrainedModel): ):
if attn_mask is not None:
features = features.masked_fill(inv_bool_mask(attn_mask), 0)
return features.sum(dim=1, keepdim=True) / attn_mask.sum(
dim=1, keepdim=True
).unsqueeze(2)
class MaxPool(nn.Module):
def forward(
self,
features: Float[Tensor, "bs seq_len feature_dim"],
attn_mask: Optional[Integer[Tensor, "bs seq_len"]] = None,
):
if attn_mask is not None:
features = features.masked_fill(inv_bool_mask(attn_mask), -float("inf"))
return torch.max(features, dim=1, keepdim=True)
class LastTokenPool(nn.Module):
def forward(
self,
features: Float[Tensor, "bs seq_len feature_dim"],
attn_mask: Optional[Integer[Tensor, "bs seq_len"]] = None,
):
last_hidden_states = (
features["hidden_states"][-1]
if "hidden_states" in features
else features["last_hidden_state"]
)
left_padding = attn_mask[:, -1].sum() == attn_mask.shape[0]
if left_padding:
return last_hidden_states[:, -1]
else:
sequence_lengths = attn_mask.sum(dim=1) - 1
batch_size = last_hidden_states.shape[0]
return last_hidden_states[
torch.arange(batch_size, device=last_hidden_states.device),
sequence_lengths,
]
AGGREGATOR = Enum("AGGREGATOR", ["MEAN", "MAX", "PERCEIVER", "LAST_TOKEN"])
AGGREGATOR_CLS = {
AGGREGATOR.MEAN: MeanPool,
AGGREGATOR.MAX: MaxPool,
AGGREGATOR.LAST_TOKEN: LastTokenPool,
AGGREGATOR.PERCEIVER: Perceiver,
}
class HyperLoRA(nn.Module):
def __init__(
self,
lora_config: LoraConfig,
aggregator: AGGREGATOR,
aggregator_kwargs: Optional[dict[str, Any]] = None,
):
super().__init__() super().__init__()
self.base_model = base_model self.lora_config = lora_config
# NOTE: we can reduce extractor_lm's depth by self.aggregator = AGGREGATOR_CLS[aggregator](**(aggregator_kwargs or {}))
# e.g., self.extractor_lm.model.layers = self.extractor_lm.model.layers[:num_extractor_layers]
# or just use the token embeddings extractor_lm.embed_tokens(input_ids)
self.extractor_lm = ...
self.hypernet = ...
def forward( def forward(
self, self,
ctx_ids: Optional[torch.LongTensor] = None, features: Float[Tensor, "bs seq_len feature_dim"],
ctx_attention_mask: Optional[torch.LongTensor] = None, attn_mask: Optional[Integer[Tensor, "bs seq_len"]] = None,
) -> Float[Tensor, "bs num_modules num_layers lora_r in_out_dim"]:
emb = self.aggregator(features, attn_mask) # [bs, n_features, feature_dim]
# loras = self.layers(emb)
# return loras
return emb
class ModulatedPretrainedModel(nn.Module):
def __init__(self, base_model: PreTrainedModel, hypernet: HyperLoRA):
super().__init__()
# self.base_model = base_model
self.device = base_model.device
# NOTE: we can reduce extractor_lm's depth by
# e.g., self.extractor_lm.model.layers = self.extractor_lm.model.layers[:num_extractor_layers]
# or just use the token embeddings extractor_lm.embed_tokens(input_ids)
# HACK: hardcode the embedding layer for now
# TODO: add explicit encoder
self.encoder = base_model.get_input_embeddings()
# register base_model as a submodule
self.register_module("base_model", base_model)
self.register_module("hypernet", hypernet)
def get_ctx_features(
self,
input_ids: Integer[Tensor, "bs seq_len"],
attention_mask: Optional[Integer[Tensor, "bs seq_len"]] = None,
):
if isinstance(self.encoder, nn.Embedding):
features = self.encoder(input_ids)
else:
features = self.encoder(input_ids=input_ids, attention_mask=attention_mask)
return features
def forward(
self,
ctx_ids: Optional[Integer[Tensor, "bs ctx_length"]] = None,
ctx_attn_mask: Optional[Integer[Tensor, "bs ctx_length"]] = None,
**model_inputs_kwargs: dict[str, Any], **model_inputs_kwargs: dict[str, Any],
) -> Union[tuple, ModelOutput]: ) -> Union[tuple, ModelOutput]:
"""Forward pass of the modulated model. """Forward pass of the modulated model."""
Args:
ctx_ids (torch.LongTensor, optional): Token IDs to be input to the hypernet
for generating LoRA parameters. Shape: (batch_size, ctx_length)
input_ids (torch.LongTensor, optional): Token IDs to be input to the base
language model. Shape: (batch_size, sequence_length)
labels (torch.LongTensor, optional): Labels for computing the language modeling
loss. Shape: (batch_size, sequence_length)
Returns:
dict: Dictionary containing:
- loss (torch.Tensor): Language modeling loss if labels are provided
- logits (torch.Tensor): Output logits from the model
"""
if ctx_ids is None: if ctx_ids is None:
logger.warning(
"No context ids provided, using the base model for the forward pass"
)
model_outputs = self.base_model(**model_inputs_kwargs) model_outputs = self.base_model(**model_inputs_kwargs)
# model_outputs.generated_loras = None # model_outputs.generated_loras = None
return model_outputs return model_outputs
loss = ... loss = ...
logits = ... logits = ...
generated_loras = ... features = self.get_ctx_features(ctx_ids, ctx_attn_mask)
generated_loras = self.hypernet(features, ctx_attn_mask)
# apply lora hook to the base model # apply lora hook to the base model
self.apply_lora_hook(generated_loras) self.apply_lora_hook(generated_loras)
@ -80,3 +163,30 @@ class ModulatedPretrainedModel(nn.Module):
def apply_lora_hook(self, generated_loras): def apply_lora_hook(self, generated_loras):
pass pass
if __name__ == "__main__":
model_name = "meta-llama/Llama-3.2-1B-Instruct"
base_model, tokenizer = get_model_and_tokenizer(
model_name,
train=True,
requires_grad=False,
peft_config=get_lora_config(model_name),
)
print(base_model)
hypernet = HyperLoRA(base_model.peft_config, aggregator=AGGREGATOR.MEAN)
model = ModulatedPretrainedModel(base_model, hypernet)
print(model)
ctx_msg = "Lorem ipsum dolor sit amet, consectetur adipiscing elit, sed do eiusmod tempor incididunt ut labore et dolore magna aliqua. Ut enim ad minim veniam, quis nostrud exercitation ullamco laboris nisi ut aliquip ex ea commodo consequat. Duis aute irure dolor in reprehenderit in voluptate velit esse cillum dolore eu fugiat nulla pariatur. Excepteur sint occaecat cupidatat non proident, sunt in culpa qui officia deserunt mollit anim id est laborum."
ctx_inputs = tokenizer(ctx_msg, return_tensors="pt").to(model.device)
ctx_features = model.get_ctx_features(**ctx_inputs)
ctx_attn_mask = ctx_inputs["attention_mask"]
print(ctx_features.shape)
agg_features = hypernet.aggregator(ctx_features, ctx_attn_mask)
print(agg_features.shape)
hnetout = hypernet(ctx_features, ctx_attn_mask)
print(hnetout.shape)
breakpoint()