mirror of
https://github.com/SakanaAI/doc-to-lora.git
synced 2026-07-26 17:11:02 +02:00
medium fw qa 3 + fallback to manual cache dataset
This commit is contained in:
parent
47edd5df1e
commit
0ff762cdff
5 changed files with 227 additions and 31 deletions
60
configs/pretrain_all_xl_3_medium.yaml
Normal file
60
configs/pretrain_all_xl_3_medium.yaml
Normal file
|
|
@ -0,0 +1,60 @@
|
||||||
|
output_dir: "" # just a placeholder
|
||||||
|
bf16: true
|
||||||
|
model_name_or_path: meta-llama/Llama-3.2-1B-Instruct
|
||||||
|
label_names: ["labels"]
|
||||||
|
# eval_on_start: True
|
||||||
|
# eval_strategy: "steps"
|
||||||
|
# eval_steps: 500
|
||||||
|
# save_strategy: "no"
|
||||||
|
# # save_steps: 500
|
||||||
|
# logging_strategy: "steps"
|
||||||
|
# logging_steps: 100
|
||||||
|
# use_liger_kernel: true
|
||||||
|
# remove_unused_columns: false
|
||||||
|
|
||||||
|
# needed to avoid OOM by compute the metrics batch by batch
|
||||||
|
# w/o this the trainer stores logits of all sample in memory...
|
||||||
|
# batch_eval_metrics: true
|
||||||
|
|
||||||
|
per_device_train_batch_size: 8
|
||||||
|
per_device_eval_batch_size: 8
|
||||||
|
max_val_samples_per_ds: 1000
|
||||||
|
# optim: schedule_free_adamw
|
||||||
|
|
||||||
|
learning_rate: 0.00004
|
||||||
|
# lr_scheduler_type: "constant_with_warmup"
|
||||||
|
neftune_noise_alpha: 5
|
||||||
|
weight_decay: 0.01
|
||||||
|
#
|
||||||
|
warmup_steps: 100
|
||||||
|
|
||||||
|
dataloader_prefetch_factor: 8
|
||||||
|
dataloader_num_workers: 8
|
||||||
|
# LoRA
|
||||||
|
lora_r: 8
|
||||||
|
lora_dropout: 0.0
|
||||||
|
target_modules:
|
||||||
|
- down_proj
|
||||||
|
# data
|
||||||
|
train_ds_names:
|
||||||
|
- fw_qa_3_medium # ~ 130M?
|
||||||
|
- ctx_qa # 300k
|
||||||
|
- pwc # 240k
|
||||||
|
- hotpot_qa # 90k
|
||||||
|
- squad # 90k
|
||||||
|
- drop # 77k
|
||||||
|
- narrativeqa # 40k
|
||||||
|
- quoref # 11k
|
||||||
|
- ropes # 11k
|
||||||
|
- synthetic_convqa # 40k
|
||||||
|
|
||||||
|
val_ds_names:
|
||||||
|
- fw_qa_3
|
||||||
|
- fw_qa_xl
|
||||||
|
- ctx_qa
|
||||||
|
- pwc
|
||||||
|
- hotpot_qa
|
||||||
|
- squad
|
||||||
|
|
||||||
|
load_best_model_at_end: true
|
||||||
|
metric_for_best_model: eval_pwc_loss
|
||||||
|
|
@ -37,7 +37,7 @@ target_modules:
|
||||||
- down_proj
|
- down_proj
|
||||||
# data
|
# data
|
||||||
train_ds_names:
|
train_ds_names:
|
||||||
- fw_qa_3_mini # ~ 267M
|
- fw_qa_3_mini # 100k
|
||||||
- ctx_qa # 300k
|
- ctx_qa # 300k
|
||||||
- pwc # 240k
|
- pwc # 240k
|
||||||
- hotpot_qa # 90k
|
- hotpot_qa # 90k
|
||||||
|
|
|
||||||
15
intx_sft.py
15
intx_sft.py
|
|
@ -3,6 +3,7 @@ import os
|
||||||
import random
|
import random
|
||||||
import string
|
import string
|
||||||
import time
|
import time
|
||||||
|
from multiprocess import set_start_method
|
||||||
from math import ceil
|
from math import ceil
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from copy import copy, deepcopy
|
from copy import copy, deepcopy
|
||||||
|
|
@ -284,7 +285,7 @@ def main():
|
||||||
training_args.lr_scheduler_type == "cosine_with_min_lr"
|
training_args.lr_scheduler_type == "cosine_with_min_lr"
|
||||||
and training_args.lr_scheduler_kwargs is None
|
and training_args.lr_scheduler_kwargs is None
|
||||||
):
|
):
|
||||||
training_args.lr_scheduler_kwargs = {"min_lr": 1e-8}
|
training_args.lr_scheduler_kwargs = {"min_lr": 1e-7}
|
||||||
args = {
|
args = {
|
||||||
**vars(deepcopy(data_args)),
|
**vars(deepcopy(data_args)),
|
||||||
**vars(deepcopy(ctx_args)),
|
**vars(deepcopy(ctx_args)),
|
||||||
|
|
@ -316,7 +317,7 @@ def main():
|
||||||
)
|
)
|
||||||
if "Llama" in ctx_name and "Vision" in ctx_name:
|
if "Llama" in ctx_name and "Vision" in ctx_name:
|
||||||
ctx_encoder_model_config = ctx_encoder_model_config.text_config
|
ctx_encoder_model_config = ctx_encoder_model_config.text_config
|
||||||
ctx_tokenizer = get_tokenizer(ctx_name)
|
ctx_tokenizer = get_tokenizer(ctx_name, train=True)
|
||||||
else:
|
else:
|
||||||
ctx_name = model.base_model.config.name_or_path
|
ctx_name = model.base_model.config.name_or_path
|
||||||
ctx_encoder_model_config = model.config
|
ctx_encoder_model_config = model.config
|
||||||
|
|
@ -353,11 +354,11 @@ def main():
|
||||||
train=True,
|
train=True,
|
||||||
use_flash_attn=model_args.use_flash_attn,
|
use_flash_attn=model_args.use_flash_attn,
|
||||||
)
|
)
|
||||||
tokenizer = get_tokenizer(model.base_model.config.name_or_path)
|
tokenizer = get_tokenizer(model.base_model.config.name_or_path, train=True)
|
||||||
ctx_name = model.ctx_encoder_args.ctx_encoder_model_name_or_path
|
ctx_name = model.ctx_encoder_args.ctx_encoder_model_name_or_path
|
||||||
if ctx_name is None:
|
if ctx_name is None:
|
||||||
ctx_name = model.base_model.config.name_or_path
|
ctx_name = model.base_model.config.name_or_path
|
||||||
ctx_tokenizer = get_tokenizer(ctx_name)
|
ctx_tokenizer = get_tokenizer(ctx_name, train=True)
|
||||||
if len([p for p in model.ctx_encoder.parameters() if p.requires_grad]):
|
if len([p for p in model.ctx_encoder.parameters() if p.requires_grad]):
|
||||||
raise ValueError("ctx_encoder contains trainable parameters")
|
raise ValueError("ctx_encoder contains trainable parameters")
|
||||||
if len([p for p in model.base_model.parameters() if p.requires_grad]):
|
if len([p for p in model.base_model.parameters() if p.requires_grad]):
|
||||||
|
|
@ -573,10 +574,12 @@ if __name__ == "__main__":
|
||||||
os.environ["WANDB_WATCH"] = "" # "all"
|
os.environ["WANDB_WATCH"] = "" # "all"
|
||||||
os.environ["WANDB_CONSOLE"] = "off"
|
os.environ["WANDB_CONSOLE"] = "off"
|
||||||
os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"
|
os.environ["PYTORCH_CUDA_ALLOC_CONF"] = "expandable_segments:True"
|
||||||
os.environ["KMP_AFFINITY"] = "disabled" # fixing iterable dataset stuck
|
# os.environ["KMP_AFFINITY"] = "disabled" # fixing iterable dataset stuck
|
||||||
os.environ["OMP_NUM_THREADS"] = "16"
|
# os.environ["OMP_NUM_THREADS"] = "16"
|
||||||
|
# os.environ["HF_DATASETS_IN_MEMORY_MAX_SIZE"] = "137438953472" # 128 TB
|
||||||
if os.getenv("DEBUG", False):
|
if os.getenv("DEBUG", False):
|
||||||
disable_caching()
|
disable_caching()
|
||||||
# randomly sleep to avoid run_name collision
|
# randomly sleep to avoid run_name collision
|
||||||
# time.sleep(random.random() * 13)
|
# time.sleep(random.random() * 13)
|
||||||
|
# set_start_method("spawn", force=True)
|
||||||
main()
|
main()
|
||||||
|
|
|
||||||
33
scripts/fw_qa_3/medium_gemma_per_rank_fac2_per_layer.sh
Normal file
33
scripts/fw_qa_3/medium_gemma_per_rank_fac2_per_layer.sh
Normal file
|
|
@ -0,0 +1,33 @@
|
||||||
|
#!/bin/bash
|
||||||
|
#SBATCH --job-name=ctxlora_medium
|
||||||
|
#SBATCH --partition=a3
|
||||||
|
#SBATCH --nodes=1
|
||||||
|
#SBATCH --gpus=8
|
||||||
|
#SBATCH --output=outputs/%x-%j.out
|
||||||
|
#SBATCH --error=outputs/%x-%j.out
|
||||||
|
|
||||||
|
# module load
|
||||||
|
# module load cuda/12.1
|
||||||
|
# module load cudnn/8.9.7
|
||||||
|
# module load nccl/cuda-12.1/2.18.3
|
||||||
|
# module load hpcx/2.20
|
||||||
|
|
||||||
|
# export OMP_NUM_THREADS=24
|
||||||
|
# export TRITON_CACHE_DIR=/tmp/.triton/
|
||||||
|
. ~/miniconda3/etc/profile.d/conda.sh
|
||||||
|
conda activate /home/rujikorn_sakana_ai/.conda/envs/ctx-to-lora
|
||||||
|
# eval "$@"
|
||||||
|
|
||||||
|
accelerate launch --num_processes=8 --gradient_accumulation_steps=4 --gradient_clipping=1.0 \
|
||||||
|
--gpu_ids all --main_process_port 29568 intx_sft.py configs/pretrain_all_xl_3_medium.yaml \
|
||||||
|
--model_name_or_path=google/gemma-2-2b-it --num_train_epochs=1 --per_device_train_batch_size=32 \
|
||||||
|
--gradient_accumulation_steps=4 --per_device_eval_batch_size=32 --exp_setup=hyper_lora --aggregator_type=perceiver \
|
||||||
|
--target_modules=down_proj \
|
||||||
|
--num_self_attends_per_block=8 --num_latent_factor=2 \
|
||||||
|
--lora_r=8 \
|
||||||
|
--eval_steps=5000 --save_steps=5000 --learning_rate=4e-5 --lora_dropout=0.0 \
|
||||||
|
--neftune_noise_alpha=5 --use_light_weight_lora=False \
|
||||||
|
--load_best_model_at_end=True --metric_for_best_model=pwc_loss --add_negative_prompt=False \
|
||||||
|
--add_repeat_prompt=False \
|
||||||
|
--use_sequence_packing=True --per_rank_gen=True \
|
||||||
|
--per_layer_processing=True
|
||||||
|
|
@ -1,9 +1,13 @@
|
||||||
import logging
|
import logging
|
||||||
from os import path
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
from os import path
|
||||||
from glob import glob
|
from glob import glob
|
||||||
from typing import Any, Callable, Iterator, Optional
|
from typing import Any, Callable, Iterator, Optional
|
||||||
|
|
||||||
|
|
||||||
|
import datasets
|
||||||
from datasets import load_dataset, IterableDataset
|
from datasets import load_dataset, IterableDataset
|
||||||
from transformers import PreTrainedTokenizerBase
|
from transformers import PreTrainedTokenizerBase
|
||||||
|
|
||||||
|
|
@ -16,10 +20,13 @@ FW_QA_PATHS = [
|
||||||
f"data/raw_datasets/fw_qa/{i:05d}.parquet" for i in [0, 1, 6, 7, 8, 10, 22, 30, 35]
|
f"data/raw_datasets/fw_qa/{i:05d}.parquet" for i in [0, 1, 6, 7, 8, 10, 22, 30, 35]
|
||||||
]
|
]
|
||||||
|
|
||||||
|
TRANSFORMED_DATA_DIR = "data/processed_datasets"
|
||||||
|
|
||||||
# approximate length of the datasets
|
# approximate length of the datasets
|
||||||
# needed for streaming datasets
|
# needed for streaming datasets
|
||||||
DS_LEN = {
|
DS_LEN = {
|
||||||
"fw_qa_3_mini": 100_000,
|
"fw_qa_3_mini": 100_000,
|
||||||
|
"fw_qa_3_medium": 121_000_000,
|
||||||
"fw_qa_3": 270_000_000,
|
"fw_qa_3": 270_000_000,
|
||||||
"fw_qa_xl": 27_000_000,
|
"fw_qa_xl": 27_000_000,
|
||||||
"ctx_qa": 300_000,
|
"ctx_qa": 300_000,
|
||||||
|
|
@ -124,6 +131,13 @@ DS_KWARGS = {
|
||||||
split="train[:100000]",
|
split="train[:100000]",
|
||||||
),
|
),
|
||||||
),
|
),
|
||||||
|
"fw_qa_3_medium": dict(
|
||||||
|
train=dict(
|
||||||
|
path="parquet",
|
||||||
|
data_files=glob("data/raw_datasets/fw_qa_3/00[0-5]*[!val].parquet"),
|
||||||
|
split="train",
|
||||||
|
),
|
||||||
|
),
|
||||||
"fw_qa_3": dict(
|
"fw_qa_3": dict(
|
||||||
train=dict(
|
train=dict(
|
||||||
path="parquet",
|
path="parquet",
|
||||||
|
|
@ -389,23 +403,14 @@ def filter_none(samples):
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
def get_tokenized_dataset(
|
def _load_and_process_dataset(
|
||||||
ds_name: str,
|
ds_name: str,
|
||||||
split: str,
|
split: str,
|
||||||
tokenizer: PreTrainedTokenizerBase,
|
|
||||||
tokenizer_kwargs: dict[str, Any],
|
|
||||||
ctx_tokenizer: PreTrainedTokenizerBase,
|
|
||||||
ctx_tokenizer_kwargs: dict[str, Any],
|
|
||||||
add_ctx_to_chat: bool,
|
|
||||||
add_repeat_prompt: bool,
|
|
||||||
add_negative_prompt: bool,
|
add_negative_prompt: bool,
|
||||||
use_kl_loss: bool,
|
add_repeat_prompt: bool,
|
||||||
set_format: Optional[str] = None,
|
streaming: bool,
|
||||||
streaming: bool = False,
|
ds_path: str,
|
||||||
) -> dict[str, Any]:
|
):
|
||||||
|
|
||||||
logger.debug(f"Loading dataset {ds_name} with split {split}...")
|
|
||||||
need_ctx_ids = not add_ctx_to_chat
|
|
||||||
try:
|
try:
|
||||||
ds_kwargs = get_ds_kwargs(ds_name, split)
|
ds_kwargs = get_ds_kwargs(ds_name, split)
|
||||||
skip = ds_kwargs.pop("skip", None)
|
skip = ds_kwargs.pop("skip", None)
|
||||||
|
|
@ -438,15 +443,74 @@ def get_tokenized_dataset(
|
||||||
cols_to_remove = [
|
cols_to_remove = [
|
||||||
col for col in ds.column_names if col not in ["context", "prompt", "response"]
|
col for col in ds.column_names if col not in ["context", "prompt", "response"]
|
||||||
]
|
]
|
||||||
ds = ds.map(get_preprocessing_fn(ds_name))
|
|
||||||
ds = ds.remove_columns(cols_to_remove)
|
ds = ds.map(
|
||||||
ds = ds.filter(filter_none, batched=True)
|
get_preprocessing_fn(ds_name),
|
||||||
ds = ds.filter(filter_long_samples, batched=True)
|
remove_columns=cols_to_remove,
|
||||||
|
num_proc=16,
|
||||||
|
)
|
||||||
|
# ds = ds.remove_columns(cols_to_remove)
|
||||||
|
ds = ds.filter(filter_none, batched=True, num_proc=16)
|
||||||
|
ds = ds.filter(filter_long_samples, batched=True, num_proc=16)
|
||||||
if split == "train":
|
if split == "train":
|
||||||
if add_negative_prompt:
|
if add_negative_prompt:
|
||||||
ds = ds.map(add_negative_prompt_fn, batched=True)
|
ds = ds.map(
|
||||||
|
add_negative_prompt_fn,
|
||||||
|
batched=True,
|
||||||
|
batch_size=100_000,
|
||||||
|
num_proc=16,
|
||||||
|
)
|
||||||
if add_repeat_prompt and "context_numbers" not in ds_name:
|
if add_repeat_prompt and "context_numbers" not in ds_name:
|
||||||
ds = ds.map(add_repeat_prompt_fn, batched=True)
|
ds = ds.map(
|
||||||
|
add_repeat_prompt_fn,
|
||||||
|
batched=True,
|
||||||
|
batch_size=100_000,
|
||||||
|
num_proc=16,
|
||||||
|
)
|
||||||
|
ds.save_to_disk(ds_path, num_proc=16)
|
||||||
|
return ds
|
||||||
|
|
||||||
|
|
||||||
|
def get_tokenized_dataset(
|
||||||
|
ds_name: str,
|
||||||
|
split: str,
|
||||||
|
tokenizer: PreTrainedTokenizerBase,
|
||||||
|
tokenizer_kwargs: dict[str, Any],
|
||||||
|
ctx_tokenizer: PreTrainedTokenizerBase,
|
||||||
|
ctx_tokenizer_kwargs: dict[str, Any],
|
||||||
|
add_ctx_to_chat: bool,
|
||||||
|
add_repeat_prompt: bool,
|
||||||
|
add_negative_prompt: bool,
|
||||||
|
use_kl_loss: bool,
|
||||||
|
set_format: Optional[str] = None,
|
||||||
|
streaming: bool = False,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
assert not use_kl_loss, "KL loss is deprecated"
|
||||||
|
logger.debug(f"Loading dataset {ds_name} with split {split}...")
|
||||||
|
need_ctx_ids = not add_ctx_to_chat
|
||||||
|
|
||||||
|
load_and_process_kwargs = dict(
|
||||||
|
ds_name=ds_name,
|
||||||
|
split=split,
|
||||||
|
add_negative_prompt=add_negative_prompt,
|
||||||
|
add_repeat_prompt=add_repeat_prompt,
|
||||||
|
streaming=streaming,
|
||||||
|
)
|
||||||
|
|
||||||
|
ds_hash = hashlib.sha256(json.dumps(load_and_process_kwargs).encode()).hexdigest()
|
||||||
|
ds_path = f"{TRANSFORMED_DATA_DIR}/{ds_hash}"
|
||||||
|
|
||||||
|
if path.exists(ds_path):
|
||||||
|
# load the cached ds
|
||||||
|
logger.info(f"Loaded processed dataset from {ds_path}")
|
||||||
|
else:
|
||||||
|
logger.info(f"Loading dataset {ds_name} with split {split}...")
|
||||||
|
_load_and_process_dataset(
|
||||||
|
**load_and_process_kwargs,
|
||||||
|
ds_path=ds_path,
|
||||||
|
)
|
||||||
|
ds = datasets.load_from_disk(ds_path)
|
||||||
|
|
||||||
tokenized_ds = construct_and_tokenize_ctx_qa(
|
tokenized_ds = construct_and_tokenize_ctx_qa(
|
||||||
tokenizer,
|
tokenizer,
|
||||||
tokenizer_kwargs,
|
tokenizer_kwargs,
|
||||||
|
|
@ -472,16 +536,47 @@ def construct_and_tokenize_ctx_qa(
|
||||||
ds,
|
ds,
|
||||||
set_format=None,
|
set_format=None,
|
||||||
):
|
):
|
||||||
|
kwargs = dict(
|
||||||
|
tokenizer=repr(tokenizer),
|
||||||
|
tokenizer_kwargs=json.dumps(tokenizer_kwargs),
|
||||||
|
ctx_tokenizer=repr(ctx_tokenizer),
|
||||||
|
ctx_tokenizer_kwargs=json.dumps(ctx_tokenizer_kwargs),
|
||||||
|
add_ctx_to_chat=add_ctx_to_chat,
|
||||||
|
use_kl_loss=use_kl_loss,
|
||||||
|
need_ctx_ids=need_ctx_ids,
|
||||||
|
ds=ds._fingerprint,
|
||||||
|
set_format=set_format,
|
||||||
|
)
|
||||||
|
kwargs_str = json.dumps(kwargs)
|
||||||
|
logger.debug(f"Tokenizing dataset with kwargs: {kwargs_str}")
|
||||||
|
ds_hash = hashlib.sha256(kwargs_str.encode()).hexdigest()
|
||||||
|
ds_path = f"{TRANSFORMED_DATA_DIR}/{ds_hash}"
|
||||||
|
if path.exists(ds_path):
|
||||||
|
# load the cached ds
|
||||||
|
logger.info(f"Loaded tokenized dataset from {ds_path}")
|
||||||
|
ds = datasets.load_from_disk(ds_path)
|
||||||
|
return ds
|
||||||
# for sft + chat_model, we need to convert the dataset to chat format
|
# for sft + chat_model, we need to convert the dataset to chat format
|
||||||
# add "messages" field
|
# add "messages" field
|
||||||
ds = ds.map(
|
ds = ds.map(
|
||||||
convert_ctx_prompt_response_to_messages,
|
convert_ctx_prompt_response_to_messages,
|
||||||
fn_kwargs={"add_ctx_to_chat": add_ctx_to_chat},
|
fn_kwargs={"add_ctx_to_chat": add_ctx_to_chat},
|
||||||
|
num_proc=16,
|
||||||
)
|
)
|
||||||
# add "chat" field
|
# add "chat" field
|
||||||
ds = ds.map(get_sft_prompt_formatting_fn(TRAINING_TASK.COMPLETION, tokenizer))
|
ds = ds.map(
|
||||||
|
get_sft_prompt_formatting_fn(TRAINING_TASK.COMPLETION, tokenizer),
|
||||||
|
batched=True,
|
||||||
|
batch_size=100_000,
|
||||||
|
num_proc=16,
|
||||||
|
)
|
||||||
# tokenize the chat + mask the assistant inputs
|
# tokenize the chat + mask the assistant inputs
|
||||||
ds = ds.filter(filter_long_chat, batched=True)
|
ds = ds.filter(
|
||||||
|
filter_long_chat,
|
||||||
|
batched=True,
|
||||||
|
batch_size=100_000,
|
||||||
|
num_proc=16,
|
||||||
|
)
|
||||||
|
|
||||||
# add "input_ids", "attention_mask", "labels"
|
# add "input_ids", "attention_mask", "labels"
|
||||||
tokenized_ds = ds.map(
|
tokenized_ds = ds.map(
|
||||||
|
|
@ -491,6 +586,7 @@ def construct_and_tokenize_ctx_qa(
|
||||||
"mask_assistant_inputs": True,
|
"mask_assistant_inputs": True,
|
||||||
"tokenizer_kwargs": tokenizer_kwargs,
|
"tokenizer_kwargs": tokenizer_kwargs,
|
||||||
},
|
},
|
||||||
|
num_proc=16,
|
||||||
)
|
)
|
||||||
|
|
||||||
# for use_kl_loss, we need "chat_ids" and "chat_attn_mask"
|
# for use_kl_loss, we need "chat_ids" and "chat_attn_mask"
|
||||||
|
|
@ -521,6 +617,8 @@ def construct_and_tokenize_ctx_qa(
|
||||||
tokenize_ctx_text,
|
tokenize_ctx_text,
|
||||||
fn_kwargs={"tokenizer": ctx_tokenizer},
|
fn_kwargs={"tokenizer": ctx_tokenizer},
|
||||||
batched=True,
|
batched=True,
|
||||||
|
batch_size=100_000,
|
||||||
|
num_proc=16,
|
||||||
)
|
)
|
||||||
|
|
||||||
tokenized_ds = tokenized_ds.remove_columns(
|
tokenized_ds = tokenized_ds.remove_columns(
|
||||||
|
|
@ -532,7 +630,9 @@ def construct_and_tokenize_ctx_qa(
|
||||||
# # the columns are unknown when using streaming dataset
|
# # the columns are unknown when using streaming dataset
|
||||||
# tokenized_ds = tokenized_ds._resolve_features()
|
# tokenized_ds = tokenized_ds._resolve_features()
|
||||||
# validate_columns(tokenized_ds)
|
# validate_columns(tokenized_ds)
|
||||||
return tokenized_ds
|
tokenized_ds.save_to_disk(ds_path, num_proc=16)
|
||||||
|
del tokenized_ds
|
||||||
|
return datasets.load_from_disk(ds_path)
|
||||||
|
|
||||||
|
|
||||||
def get_sft_prompt_formatting_fn(
|
def get_sft_prompt_formatting_fn(
|
||||||
|
|
|
||||||
Loading…
Add table
Add a link
Reference in a new issue