diff --git a/configs/context_numbers_10.yaml b/configs/context_numbers_10.yaml index ba26c4c..d29b9ee 100644 --- a/configs/context_numbers_10.yaml +++ b/configs/context_numbers_10.yaml @@ -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 diff --git a/configs/context_numbers_128.yaml b/configs/context_numbers_128.yaml index 6faf010..3d5e073 100644 --- a/configs/context_numbers_128.yaml +++ b/configs/context_numbers_128.yaml @@ -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 diff --git a/data/raw_datasets/generate_data.py b/data/raw_datasets/generate_data.py index f950dc5..5c788c3 100644 --- a/data/raw_datasets/generate_data.py +++ b/data/raw_datasets/generate_data.py @@ -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}") \ No newline at end of file + print(f"Dataset generated and saved at {save_dir}") diff --git a/hyperlora/configs.py b/hyperlora/configs.py index d075498..1c9296b 100644 --- a/hyperlora/configs.py +++ b/hyperlora/configs.py @@ -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) diff --git a/hyperlora/intx_sft.py b/hyperlora/intx_sft.py index 64bc1b8..a8760b0 100644 --- a/hyperlora/intx_sft.py +++ b/hyperlora/intx_sft.py @@ -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 diff --git a/hyperlora/modeling_utils.py b/hyperlora/modeling_utils.py index 9f94847..ec226ad 100644 --- a/hyperlora/modeling_utils.py +++ b/hyperlora/modeling_utils.py @@ -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) diff --git a/hyperlora/training_utils.py b/hyperlora/training_utils.py index 834bfbd..142533d 100644 --- a/hyperlora/training_utils.py +++ b/hyperlora/training_utils.py @@ -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) diff --git a/hyperlora/utils.py b/hyperlora/utils.py index 73a2780..223405b 100644 --- a/hyperlora/utils.py +++ b/hyperlora/utils.py @@ -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