mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-26 17:11:02 +02:00
per-rank bias + skip interleave if cached+ add bnb
This commit is contained in:
parent
d6c649e13c
commit
178e7165ea
12 changed files with 302 additions and 160 deletions
47
train.py
47
train.py
|
|
@ -3,15 +3,11 @@ import logging
|
|||
import os
|
||||
from copy import deepcopy
|
||||
from functools import partial
|
||||
from math import isclose
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import wandb
|
||||
from datasets import (
|
||||
disable_caching,
|
||||
interleave_datasets,
|
||||
)
|
||||
from datasets import disable_caching
|
||||
from peft import PeftModel
|
||||
from transformers import (
|
||||
AutoConfig,
|
||||
|
|
@ -66,24 +62,6 @@ logger = logging.getLogger()
|
|||
LOCAL_RANK = int(os.getenv("LOCAL_RANK", "0"))
|
||||
|
||||
|
||||
def get_ds_prob(train_ds_len: list[int], total_len: int):
|
||||
# if a dataset is smaller than 1%, make it 1%
|
||||
probs = [0 for _ in train_ds_len]
|
||||
for i, ds_len in enumerate(train_ds_len):
|
||||
if ds_len / total_len <= 0.01:
|
||||
probs[i] = 0.01
|
||||
res_probs = 1 - sum(probs)
|
||||
res_total_len = sum([l for l in train_ds_len if (l / total_len) > 0.01])
|
||||
for i, ds_len in enumerate(train_ds_len):
|
||||
if (ds_len / total_len) > 0.01:
|
||||
probs[i] = ds_len / res_total_len * res_probs
|
||||
logger.debug(f"Dataset probabilities: {probs}")
|
||||
assert isclose(sum(probs), 1.0), (
|
||||
f"Probs sum to {sum(probs)} ({probs}), expected 1.0"
|
||||
)
|
||||
return probs
|
||||
|
||||
|
||||
def main():
|
||||
############ Argument parsing
|
||||
parser = ArgumentParser(
|
||||
|
|
@ -183,6 +161,7 @@ def main():
|
|||
train=True,
|
||||
requires_grad=False, # ctx_args.exp_setup == ExperimentSetup.FULL_FINETUNE,
|
||||
peft_config=get_lora_config(model_name, **vars(lora_args)),
|
||||
# use_q_lora=True,
|
||||
)
|
||||
ctx_name = ctx_encoder_args.ctx_encoder_model_name_or_path
|
||||
if ctx_name is not None:
|
||||
|
|
@ -328,31 +307,21 @@ def main():
|
|||
val_indices = np.random.permutation(len(ds))[:n_val_samples]
|
||||
val_ds[ds_name] = val_ds[ds_name].select(val_indices)
|
||||
|
||||
train_ds_len = [len(ds) for ds in train_ds.values()]
|
||||
total_len = sum(train_ds_len)
|
||||
|
||||
# TODO: cache this?
|
||||
train_ds = interleave_datasets(
|
||||
list(train_ds.values()),
|
||||
probabilities=get_ds_prob(train_ds_len, total_len),
|
||||
seed=training_args.seed,
|
||||
stopping_strategy="all_exhausted",
|
||||
)
|
||||
logger.info(f"Train dataset length: {len(train_ds)}")
|
||||
|
||||
total_samples = sum([len(ds) for ds in train_ds.values()])
|
||||
logging.info(f"Total samples before packing: {total_samples}")
|
||||
logging.info("Packing dataset")
|
||||
with training_args.main_process_first():
|
||||
logging.info("Packing dataset")
|
||||
old_ds_len = len(train_ds)
|
||||
train_ds = pack(
|
||||
train_ds,
|
||||
ctx_args.max_packed_inp_len,
|
||||
ctx_args.max_packed_ctx_len,
|
||||
max_packed_size=-1,
|
||||
seed=training_args.seed,
|
||||
num_proc=16,
|
||||
)
|
||||
logger.info(f"Train dataset length: {len(train_ds)}")
|
||||
logger.info(f"Packed dataset length: {len(train_ds)}")
|
||||
logger.info(
|
||||
f"Avg. # of samples per packed sequence: {old_ds_len / len(train_ds)}"
|
||||
f"Avg. # of samples per packed sequence: {total_samples / len(train_ds)}"
|
||||
)
|
||||
logger.info("Setting per_device_train_batch_size to 1")
|
||||
training_args.per_device_train_batch_size = 1
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue