mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-23 17:01:04 +02:00
biashyperinit + better config + bigger ctx_num_10 + perceiver args
This commit is contained in:
parent
d83c1ccc6d
commit
9641f5c44a
8 changed files with 218 additions and 91 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue