biashyperinit + better config + bigger ctx_num_10 + perceiver args

This commit is contained in:
51616 2025-01-04 15:42:08 +00:00
parent d83c1ccc6d
commit 9641f5c44a
8 changed files with 218 additions and 91 deletions

View file

@ -22,7 +22,7 @@ per_device_eval_batch_size: 128
learning_rate: 0.00001
# lr_scheduler_type: "constant_with_warmup"
neftune_noise_alpha: 1
weight_decay: 0.01
weight_decay: 0.1
warmup_ratio: 0.1
# LoRA
@ -35,26 +35,26 @@ target_modules:
# data
train_ds_names:
- data/raw_datasets/context_numbers_2
- data/raw_datasets/context_numbers_3
- data/raw_datasets/context_numbers_4
- data/raw_datasets/context_numbers_5
- data/raw_datasets/context_numbers_6
- data/raw_datasets/context_numbers_7
- data/raw_datasets/context_numbers_8
- data/raw_datasets/context_numbers_9
- data/raw_datasets/context_numbers_10
- data/raw_datasets/context_numbers_2_10k
- data/raw_datasets/context_numbers_3_10k
- data/raw_datasets/context_numbers_4_10k
- data/raw_datasets/context_numbers_5_10k
- data/raw_datasets/context_numbers_6_10k
- data/raw_datasets/context_numbers_7_10k
- data/raw_datasets/context_numbers_8_10k
- data/raw_datasets/context_numbers_9_10k
- data/raw_datasets/context_numbers_10_10k
val_ds_names:
- data/raw_datasets/context_numbers_2
- data/raw_datasets/context_numbers_3
- data/raw_datasets/context_numbers_4
- data/raw_datasets/context_numbers_5
- data/raw_datasets/context_numbers_6
- data/raw_datasets/context_numbers_7
- data/raw_datasets/context_numbers_8
- data/raw_datasets/context_numbers_9
- data/raw_datasets/context_numbers_10
- data/raw_datasets/context_numbers_2_10k
- data/raw_datasets/context_numbers_3_10k
- data/raw_datasets/context_numbers_4_10k
- data/raw_datasets/context_numbers_5_10k
- data/raw_datasets/context_numbers_6_10k
- data/raw_datasets/context_numbers_7_10k
- data/raw_datasets/context_numbers_8_10k
- data/raw_datasets/context_numbers_9_10k
- data/raw_datasets/context_numbers_10_10k
test_ds_names:
- data/raw_datasets/context_numbers_11

View file

@ -22,7 +22,7 @@ per_device_eval_batch_size: 128
learning_rate: 0.00001
# lr_scheduler_type: "constant_with_warmup"
neftune_noise_alpha: 1
weight_decay: 0.01
weight_decay: 0.1
warmup_ratio: 0.1
# LoRA

View file

@ -24,7 +24,9 @@ def get_random_combinations(numbers: list[int], n: int, k: int) -> list[list[int
return zip(*[[numbers[i] for i in ind] for ind in indices])
def generate_number_dataset(max_num: int = 1000, n: int = 12000, k: int = 3, save_dir: str = None):
def generate_number_dataset(
max_num: int = 1000, n: int = 12000, k: int = 3, save_dir: str = None
):
"""
Generate a dataset of numbers with corresponding query and answer,
split into train/val/test sets.
@ -62,7 +64,7 @@ def generate_number_dataset(max_num: int = 1000, n: int = 12000, k: int = 3, sav
train_data = dataset[:train_size]
val_data = dataset[train_size : train_size + val_size]
test_data = dataset[train_size + val_size :]
save_dir = "" if save_dir is None else save_dir
# Save splits to separate files
save_jsonl(train_data, f"{save_dir}/train.jsonl")
@ -75,23 +77,30 @@ if __name__ == "__main__":
random.seed(42)
# Generate dataset
for k in range(2, 129):
for k in range(2, 11):
save_dir = f"context_numbers_{k}_10k"
os.makedirs(save_dir, exist_ok=True)
generate_number_dataset(n=12000, k=k, save_dir=save_dir)
print(f"Dataset generated and saved at {save_dir}")
for k in range(12, 129):
save_dir = f"context_numbers_{k}"
os.makedirs(save_dir, exist_ok=True)
generate_number_dataset(n=1200, k=k, save_dir=save_dir)
print(f"Dataset generated and saved at {save_dir}")
for k in range(144,257,16):
for k in range(144, 257, 16):
save_dir = f"context_numbers_{k}"
os.makedirs(save_dir, exist_ok=True)
generate_number_dataset(n=12000, k=k, save_dir=save_dir)
print(f"Dataset generated and saved at {save_dir}")
for k in [512, 1024, 2048]:
save_dir = f"context_numbers_{k}"
os.makedirs(save_dir, exist_ok=True)
generate_number_dataset(n=12000, k=k, save_dir=save_dir)
print(f"Dataset generated and saved at {save_dir}")
print(f"Dataset generated and saved at {save_dir}")

View file

@ -6,7 +6,6 @@ from enum import Enum, auto
from typing import Any, Dict, List, Literal, NewType, Optional, Tuple
import yaml
from modeling_utils import AGGREGATOR_TYPE
from transformers import MODEL_FOR_CAUSAL_LM_MAPPING, HfArgumentParser
MODEL_CONFIG_CLASSES = list(MODEL_FOR_CAUSAL_LM_MAPPING.keys())
@ -162,10 +161,6 @@ class CtxTrainingArguments:
default=2**13,
metadata={"help": "Maximum base length for training."},
)
aggregator_type: AGGREGATOR_TYPE = field(
default=AGGREGATOR_TYPE.POOLER,
metadata={"help": "Aggregator type for HyperLoRA."},
)
@dataclass
@ -184,6 +179,55 @@ class DataArguments:
)
@dataclass
class HypernetArguments:
latent_size: int = field(
default=512,
metadata={"help": "Latent size for HyperLoRA."},
)
@dataclass
class AggregatorArguments:
aggregator_type: Literal["pooler", "perceiver"] = field(
default="pooler",
metadata={"help": "Aggregator type for HyperLoRA."},
)
# pooler
pooling_type: str = field(
default="mean",
metadata={"help": "Pooling type for HyperLoRA."},
)
# feature_size: int
# num_layers: int
# num_modules: int
# output_size: int
# perceiver
attention_probs_dropout_prob: float = field(
default=0.0,
metadata={"help": "Attention dropout probability for Perceiver."},
)
num_blocks: int = field(
default=8,
metadata={"help": "Number of blocks for Perceiver."},
)
num_self_attends_per_block: int = field(
default=6,
metadata={"help": "Number of self-attends per block for Perceiver."},
)
self_attention_widening_factor: int = field(
default=1,
metadata={"help": "Self-attention widening factor for Perceiver."},
)
cross_attention_widening_factor: int = field(
default=1,
metadata={"help": "Cross-attention widening factor for Perceiver."},
)
if __name__ == "__main__":
print(ExperimentSetup)
print(ExperimentSetup.LORA)

View file

@ -55,6 +55,8 @@ from configs import (
ExperimentSetup,
LoRAArguments,
ModelArguments,
HypernetArguments,
AggregatorArguments,
)
logger = logging.getLogger()
@ -173,12 +175,32 @@ def main(output_dir):
ModelArguments,
LoRAArguments,
TrainingArguments,
HypernetArguments,
AggregatorArguments,
)
)
data_args, ctx_args, model_args, lora_args, training_args = parser.parse()
(
data_args,
ctx_args,
model_args,
lora_args,
training_args,
hypernet_args,
aggregator_args,
) = parser.parse()
# there shouldn't be overlap between args
validate_args([data_args, ctx_args, model_args, lora_args, training_args])
validate_args(
[
data_args,
ctx_args,
model_args,
lora_args,
training_args,
hypernet_args,
aggregator_args,
]
)
args = {
**vars(data_args),
@ -186,17 +208,14 @@ def main(output_dir):
**vars(model_args),
**vars(lora_args),
**vars(training_args),
**vars(hypernet_args),
**vars(aggregator_args),
}
run_name = os.path.basename(output_dir)
training_args.run_name = run_name
training_args.output_dir = output_dir
training_args.logging_dir = output_dir
logger.info(f"run_name: {run_name}")
logger.info(f"data_args: {data_args}")
logger.info(f"ctx_args: {ctx_args}")
logger.info(f"model_args: {model_args}")
logger.info(f"lora_args: {lora_args}")
logger.debug(f"args: {args}")
############ Model setup
@ -212,7 +231,8 @@ def main(output_dir):
if ctx_args.exp_setup == ExperimentSetup.HYPER_LORA:
logger.info("Using HyperLoRA")
hypernet = HyperLoRA(
get_hypernet_config(model, aggregator_type=ctx_args.aggregator_type)
get_hypernet_config(model, hypernet_args, aggregator_args),
model,
).to(model.device)
# HACK: hardcode the embedding layer for now
# TODO: add ctx_encoder config

View file

@ -8,13 +8,17 @@ from typing import Any, Iterable, Optional, Tuple, Union
import torch
import torch.nn.functional as F
from configs import AggregatorArguments, HypernetArguments
from einops import rearrange, repeat, unpack
from einops.layers.torch import EinMix as Mix
from einops.layers.torch import Reduce
from hooks import add_generated_lora_hook, remove_hook_handles
from jaxtyping import Float, Integer
from model_loading import get_lora_config, get_model_and_tokenizer
from peft import LoraConfig
from peft import get_peft_config, load_peft_weights, LoraConfig, PeftConfig, PeftModel
from peft.tuners._buffer_dict import BufferDict
from peft.tuners.tuners_utils import BaseTunerLayer, check_target_module_exists
from pooling import POOL_FN, get_pooling_fn
from torch import Tensor, einsum, nn
from transformers import PerceiverConfig, PerceiverModel, PreTrainedModel
@ -34,6 +38,8 @@ class AGGREGATOR_TYPE(str, Enum):
@dataclass
class AggregatorConfig:
aggregator_type: AGGREGATOR_TYPE
# pooler
pooling_type: POOL_FN
feature_size: int
@ -41,41 +47,24 @@ class AggregatorConfig:
num_modules: int
output_size: int
# # perceiver
# depth: int = 8
# input_channels: int = 2048
# input_axis: int = 1
# num_latents: int = 512
# latent_dim: int = 512
# cross_heads: int = 1
# latent_heads: int = 8
# cross_dim_head: int = 64
# latent_dim_head: int = 64
# attn_dropout: float = 0.0
# ff_dropout: float = 0.0
# weight_tie_layers: bool = False
# self_per_cross_attn: int = 1
# final_classifier_head: bool = False
# num_classes: int = 1000
# fourier_encode_data: bool = False
# num_freq_bands: int = 16
# max_freq: float = 10.0
# perceiver
attention_probs_dropout_prob: float = 0.0
num_blocks: int = 1
num_self_attends_per_block: int = 16
self_attention_widening_factor: int = 4
cross_attention_widening_factor: int = 1
def get_aggregator_config(
model: PreTrainedModel,
output_size: int,
pooling_type: POOL_FN = POOL_FN.MEAN,
model: PreTrainedModel, output_size: int, aggregator_args: AggregatorArguments
):
lora_config = model.peft_config["default"]
return AggregatorConfig(
pooling_type=pooling_type,
feature_size=model.config.hidden_size,
output_size=output_size,
num_layers=get_num_layers(model),
num_modules=len(lora_config.target_modules),
**vars(aggregator_args),
)
@ -86,25 +75,27 @@ class HypernetConfig:
module_names: dict[str, list[str]]
layer_indices: Iterable[int]
feature_sizes: tuple[dict[str, int], dict[str, int]]
aggregator_type: AGGREGATOR_TYPE
aggregator_config: AggregatorConfig
def get_hypernet_config(
model: PreTrainedModel,
latent_size: int = 512,
aggregator_type: AGGREGATOR_TYPE = AGGREGATOR_TYPE.POOLER,
hypernet_args: HypernetArguments,
aggregator_args: AggregatorArguments,
):
lora_config = model.peft_config["default"]
indices = torch.arange(get_num_layers(model), device=model.device)
return HypernetConfig(
latent_size=latent_size,
latent_size=hypernet_args.latent_size,
lora_config=lora_config,
module_names=get_lora_module_names(model, lora_config.target_modules, indices),
layer_indices=indices,
feature_sizes=get_peft_in_out_features(model, peft_config=lora_config),
aggregator_type=aggregator_type,
aggregator_config=get_aggregator_config(model, latent_size, POOL_FN.MEAN),
aggregator_config=get_aggregator_config(
model,
hypernet_args.latent_size,
aggregator_args,
),
)
@ -121,9 +112,10 @@ class Perceiver(nn.Module):
d_model=feature_size, # + num_bands
num_latents=num_layers * num_modules * 8,
d_latents=output_size,
attention_probs_dropout_prob=0.0,
num_blocks=1,
num_self_attends_per_block=26,
# attention_probs_dropout_prob=0.0,
# num_blocks=8,
# num_self_attends_per_block=6,
# self_attention_widening_factor=4,
**kwargs,
)
decoder = PerceiverBasicDecoder(
@ -331,10 +323,41 @@ class EarlyExit(nn.Module):
return model_outputs.last_hidden_state
def init_mixer_weights(m: nn.Module):
# bias-hyperinit
# init weights to zeros and bias to the base weights
...
def get_init_peft_weights(model: PeftModel, peft_config: PeftConfig = None):
if peft_config is None:
peft_config = model.peft_config["default"]
peft_weights = {module_name: dict() for module_name in peft_config.target_modules}
adapter_name = "default"
for module_name, module in model.named_modules():
if not check_target_module_exists(peft_config, module_name):
continue
if not isinstance(module, BaseTunerLayer):
continue
# support just Linear layer for now
# all modules should be a leave module that is Linear layer
assert isinstance(
module.base_layer, nn.Linear
), "all modules should be a leave module that is Linear layer"
# this should always pass
name = module_name.split(".")[-1]
assert name in peft_config.target_modules
for submodule_name, submodule in module.named_modules():
if not isinstance(submodule, (nn.ModuleDict, nn.ParameterDict, BufferDict)):
continue
if adapter_name not in submodule:
continue
if submodule_name not in peft_weights[name]:
peft_weights[name][submodule_name] = submodule[adapter_name]
else:
smod1 = peft_weights[name][submodule_name]
smod2 = submodule[adapter_name]
assert type(smod1) == type(smod2)
return peft_weights
class HyperLoRA(nn.Module):
@ -347,17 +370,15 @@ class HyperLoRA(nn.Module):
# aggregator_type: AGGREGATOR_TYPE,
# aggregator_kwargs: dict,
config: HypernetConfig,
base_model: Optional[PreTrainedModel] = None,
):
super().__init__()
# NOTE: this class then only handles the output space of the hypernet
# e.g., shared_AB_head, per_rank_gen, etc.
# TODO: add different output spaces
# aggregator output [bs, n_layers, n_modules, feature_dim]
# by mixing the pooled features with layer embs and module embs (for pooling)
# or via a perceiver w/ bottleneck size = n_modules * n_layers
self.aggregator = AG[config.aggregator_type](**vars(config.aggregator_config))
agg_config = config.aggregator_config
self.aggregator = AG[agg_config.aggregator_type](**vars(agg_config))
self.lora_config = config.lora_config
@ -382,12 +403,33 @@ class HyperLoRA(nn.Module):
self.head = Mix(
"bs n_layers n_modules d -> bs n_layers n_modules r out_d",
weight_shape="n_modules d r out_d",
bias_shape=None, # no bias
# bias_shape=None, # no bias
bias_shape="n_modules r out_d",
n_modules=len(self.target_modules),
d=config.latent_size,
r=config.lora_config.r,
out_d=max(self.in_d[m] + self.out_d[m] for m in self.target_modules),
)
if base_model is not None:
self._init_head(base_model)
@torch.no_grad()
def _init_head(self, base_model: PreTrainedModel):
peft_weights = get_init_peft_weights(base_model, self.lora_config)
logger.debug(f"peft_weights: {peft_weights}")
self.head.weight.data[:] = 0
self.head.bias.data[:] = 0
for i, m in enumerate(self.target_modules):
A = peft_weights[m]["lora_A"].weight.clone() # [r, in_d]
B = peft_weights[m]["lora_B"].weight.clone() # [out_d, r]
biases = [A, B.T]
# bias_hyper_init(self.head.bias[i], biases)
# bias-hyperinit
# init weights to zeros and bias to the base weights
bias_cat = torch.cat(biases, dim=1)
self.head.bias.data[..., i, :, : bias_cat.shape[1]] = bias_cat
self.head.bias.requires_grad = False
def _to_lora_dict(
self, flat_loras: Float[Tensor, "bs n_layers n_modules r max_io_dim"]
@ -628,6 +670,7 @@ if __name__ == "__main__":
peft_config=get_lora_config(model_name),
)
print(base_model)
device = base_model.device
# lora_config = base_model.peft_config["default"]
# in_d, out_d = get_peft_in_out_features(base_model, peft_config=lora_config)
@ -640,15 +683,13 @@ if __name__ == "__main__":
# ).to(base_model.device)
ctx_encoder = EarlyExit(get_base_model(base_model), 4)
hypernet = HyperLoRA(get_hypernet_config(base_model)).to(base_model.device)
hypernet = HyperLoRA(get_hypernet_config(base_model), base_model).to(device)
model = ModulatedPretrainedModel(base_model, hypernet, ctx_encoder).to(
base_model.device
)
model = ModulatedPretrainedModel(base_model, hypernet, ctx_encoder).to(device)
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_inputs = tokenizer(ctx_msg, return_tensors="pt").to(device)
ctx_ids = ctx_inputs["input_ids"]
ctx_attn_mask = ctx_inputs["attention_mask"]
ctx_features = model.ctx_encoder(input_ids=ctx_ids, attention_mask=ctx_attn_mask)
@ -664,7 +705,7 @@ if __name__ == "__main__":
print(hnetout.shape)
prompt_msg = "hello"
prompt_inputs = tokenizer(prompt_msg, return_tensors="pt").to(model.device)
prompt_inputs = tokenizer(prompt_msg, return_tensors="pt").to(device)
basemodelout = model.base_model(**prompt_inputs)
print(basemodelout)

View file

@ -1,3 +1,4 @@
import gc
import json
import logging
from collections import defaultdict
@ -6,6 +7,7 @@ from enum import Enum
import numpy as np
from rouge_score import rouge_scorer
import torch
from transformers import (
GenerationConfig,
Seq2SeqTrainer,
@ -19,6 +21,13 @@ TRAINING_TASK = Enum("TRAINING_TASK", ["CAUSAL_LM", "COMPLETION"])
logger = logging.getLogger()
def clear_gpu():
gc.collect()
torch.cuda.empty_cache()
torch.cuda.reset_max_memory_allocated()
torch.cuda.reset_max_memory_cached()
def compute_rouge(pred_texts, label_texts):
out = defaultdict(list)
scorer = rouge_scorer.RougeScorer(["rouge1", "rougeL"], use_stemmer=False)
@ -170,6 +179,8 @@ def train_model(
trainer.save_metrics("eval", metrics)
trainer.save_model()
clear_gpu()
############## Evaluation
# TODO: generalize gen_kwargs for validation
max_new_tokens = 100
@ -220,6 +231,7 @@ def train_model(
if isinstance(val_dataset, dict):
val_dataset = val_dataset["val"]
eval_generation(eval_trainer, tokenizer, val_dataset, "val", gen_kwargs)
clear_gpu()
if test_dataset is not None:
eval_generation(eval_trainer, tokenizer, test_dataset, "test", gen_kwargs)

View file

@ -131,6 +131,7 @@ def validate_args(args_list):
# there shouldn't be overlap between args
keys = set()
for args in args_list:
logger.debug(args)
args_keys = set(vars(args).keys())
assert len(keys & args_keys) == 0, "Overlap between args"
keys |= args_keys