per-rank bias + skip interleave if cached+ add bnb

This commit is contained in:
51616 2025-08-09 15:20:37 +00:00
parent d6c649e13c
commit 178e7165ea
12 changed files with 302 additions and 160 deletions

View file

@ -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