RexTRO111's picture
Upload 3 files
e6cd8c3 verified
Raw
History Blame Contribute Delete
10.8 kB
import modal
APP_NAME = "pythia-410m-dolly-lora"
VOLUME_NAME = "pythia-410m-dolly-output"
MODEL_NAME = "EleutherAI/pythia-410m"
DATASET_NAME = "databricks/databricks-dolly-15k"
VOLUME_ROOT = "/data"
OUTPUT_DIR = f"{VOLUME_ROOT}/outputs/pythia-410m-dolly-lora"
CACHE_DIR = f"{VOLUME_ROOT}/cache/huggingface"
app = modal.App(APP_NAME)
output_volume = modal.Volume.from_name(
VOLUME_NAME,
create_if_missing=True,
)
image = (
modal.Image.debian_slim(python_version="3.11")
.pip_install(
"torch==2.7.1",
"transformers==4.53.2",
"datasets==3.6.0",
"peft==0.16.0",
"accelerate==1.8.1",
"safetensors==0.5.3",
"sentencepiece==0.2.0",
)
.env(
{
"HF_HOME": CACHE_DIR,
"TRANSFORMERS_CACHE": CACHE_DIR,
"HF_DATASETS_CACHE": f"{CACHE_DIR}/datasets",
"TOKENIZERS_PARALLELISM": "false",
}
)
)
@app.function(
image=image,
gpu="A10",
cpu=4.0,
memory=16384,
timeout=60 * 60 * 12,
volumes={
VOLUME_ROOT: output_volume,
},
)
def train():
import json
import os
import time
import torch
from datasets import load_dataset
from peft import LoraConfig, TaskType, get_peft_model
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
Trainer,
TrainingArguments,
set_seed,
)
# ---------------------------------------------------------
# Configuration
# ---------------------------------------------------------
seed = 3407
max_length = 768
num_train_epochs = 3
learning_rate = 2e-4
per_device_train_batch_size = 4
per_device_eval_batch_size = 4
gradient_accumulation_steps = 4
lora_rank = 16
lora_alpha = 32
lora_dropout = 0.05
set_seed(seed)
os.makedirs(OUTPUT_DIR, exist_ok=True)
print("=" * 72)
print("Pythia-410M Dolly 15K FP16 LoRA")
print(f"Model: {MODEL_NAME}")
print(f"Dataset: {DATASET_NAME}")
print(f"Max length: {max_length}")
print(f"Output: {OUTPUT_DIR}")
print("=" * 72)
if not torch.cuda.is_available():
raise RuntimeError("CUDA GPU was not detected.")
gpu_name = torch.cuda.get_device_name(0)
print(f"GPU: {gpu_name}")
print(
f"GPU memory: "
f"{torch.cuda.get_device_properties(0).total_memory / 1024**3:.2f} GiB"
)
# ---------------------------------------------------------
# Tokenizer
# ---------------------------------------------------------
tokenizer = AutoTokenizer.from_pretrained(
MODEL_NAME,
cache_dir=CACHE_DIR,
use_fast=True,
)
# Pythia's tokenizer does not provide a separate padding token.
# Reusing EOS avoids increasing the vocabulary.
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "right"
tokenizer.model_max_length = max_length
# ---------------------------------------------------------
# Dataset
# ---------------------------------------------------------
raw_dataset = load_dataset(
DATASET_NAME,
split="train",
cache_dir=f"{CACHE_DIR}/datasets",
)
split_dataset = raw_dataset.train_test_split(
test_size=0.02,
seed=seed,
shuffle=True,
)
train_dataset = split_dataset["train"]
eval_dataset = split_dataset["test"]
def build_prompt(example: dict) -> tuple[str, str]:
instruction = str(example.get("instruction", "")).strip()
context = str(example.get("context", "") or "").strip()
response = str(example.get("response", "")).strip()
if context:
prompt = (
"### Instruction:\n"
f"{instruction}\n\n"
"### Context:\n"
f"{context}\n\n"
"### Response:\n"
)
else:
prompt = (
"### Instruction:\n"
f"{instruction}\n\n"
"### Response:\n"
)
answer = response + tokenizer.eos_token
return prompt, answer
def tokenize_example(example: dict) -> dict:
prompt, answer = build_prompt(example)
# Tokenize separately so labels can mask the instruction.
prompt_ids = tokenizer(
prompt,
add_special_tokens=False,
truncation=False,
)["input_ids"]
answer_ids = tokenizer(
answer,
add_special_tokens=False,
truncation=False,
)["input_ids"]
# Keep the response whenever possible.
if len(answer_ids) >= max_length:
answer_ids = answer_ids[:max_length]
prompt_ids = []
else:
maximum_prompt_tokens = max_length - len(answer_ids)
prompt_ids = prompt_ids[-maximum_prompt_tokens:]
input_ids = prompt_ids + answer_ids
labels = [-100] * len(prompt_ids) + answer_ids.copy()
attention_mask = [1] * len(input_ids)
padding_length = max_length - len(input_ids)
input_ids += [tokenizer.pad_token_id] * padding_length
attention_mask += [0] * padding_length
labels += [-100] * padding_length
return {
"input_ids": input_ids,
"attention_mask": attention_mask,
"labels": labels,
}
original_columns = train_dataset.column_names
train_dataset = train_dataset.map(
tokenize_example,
remove_columns=original_columns,
desc="Tokenizing training dataset",
num_proc=4,
)
eval_dataset = eval_dataset.map(
tokenize_example,
remove_columns=original_columns,
desc="Tokenizing evaluation dataset",
num_proc=4,
)
train_dataset.set_format(type="torch")
eval_dataset.set_format(type="torch")
print(f"Training examples: {len(train_dataset):,}")
print(f"Evaluation examples: {len(eval_dataset):,}")
# ---------------------------------------------------------
# Model
# ---------------------------------------------------------
model = AutoModelForCausalLM.from_pretrained(
MODEL_NAME,
cache_dir=CACHE_DIR,
torch_dtype=torch.float16,
low_cpu_mem_usage=True,
)
model.config.pad_token_id = tokenizer.pad_token_id
model.config.use_cache = False
# Saves activation memory at the cost of some extra computation.
model.gradient_checkpointing_enable()
# GPT-NeoX attention projection names used by Pythia.
lora_config = LoraConfig(
task_type=TaskType.CAUSAL_LM,
inference_mode=False,
r=lora_rank,
lora_alpha=lora_alpha,
lora_dropout=lora_dropout,
target_modules=[
"query_key_value",
"dense",
],
bias="none",
)
model = get_peft_model(model, lora_config)
model.enable_input_require_grads()
model.print_trainable_parameters()
# ---------------------------------------------------------
# Training
# ---------------------------------------------------------
training_args = TrainingArguments(
output_dir=OUTPUT_DIR,
num_train_epochs=num_train_epochs,
learning_rate=learning_rate,
per_device_train_batch_size=per_device_train_batch_size,
per_device_eval_batch_size=per_device_eval_batch_size,
gradient_accumulation_steps=gradient_accumulation_steps,
fp16=True,
bf16=False,
tf32=True,
optim="adamw_torch_fused",
lr_scheduler_type="cosine",
warmup_ratio=0.03,
weight_decay=0.01,
max_grad_norm=1.0,
logging_strategy="steps",
logging_steps=20,
eval_strategy="steps",
eval_steps=100,
save_strategy="steps",
save_steps=100,
save_total_limit=3,
load_best_model_at_end=True,
metric_for_best_model="eval_loss",
greater_is_better=False,
dataloader_num_workers=4,
dataloader_pin_memory=True,
report_to="none",
remove_unused_columns=False,
seed=seed,
data_seed=seed,
# Resume safely after an interruption.
save_safetensors=True,
)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=train_dataset,
eval_dataset=eval_dataset,
processing_class=tokenizer,
)
checkpoints = []
if os.path.isdir(OUTPUT_DIR):
checkpoints = sorted(
[
os.path.join(OUTPUT_DIR, name)
for name in os.listdir(OUTPUT_DIR)
if name.startswith("checkpoint-")
and os.path.isdir(os.path.join(OUTPUT_DIR, name))
],
key=lambda path: int(path.rsplit("-", 1)[-1]),
)
resume_checkpoint = checkpoints[-1] if checkpoints else None
if resume_checkpoint:
print(f"Resuming from checkpoint: {resume_checkpoint}")
else:
print("Starting a fresh training run.")
start_time = time.time()
train_result = trainer.train(
resume_from_checkpoint=resume_checkpoint,
)
elapsed_seconds = time.time() - start_time
# ---------------------------------------------------------
# Save final LoRA adapter and tokenizer
# ---------------------------------------------------------
final_adapter_dir = os.path.join(OUTPUT_DIR, "final-adapter")
trainer.model.save_pretrained(
final_adapter_dir,
safe_serialization=True,
)
tokenizer.save_pretrained(final_adapter_dir)
trainer.save_state()
metrics = dict(train_result.metrics)
metrics["elapsed_seconds_measured"] = elapsed_seconds
metrics["base_model"] = MODEL_NAME
metrics["dataset"] = DATASET_NAME
metrics["max_length"] = max_length
metrics["gpu"] = gpu_name
metrics["lora_rank"] = lora_rank
metrics["lora_alpha"] = lora_alpha
metrics["train_examples"] = len(train_dataset)
metrics["eval_examples"] = len(eval_dataset)
with open(
os.path.join(OUTPUT_DIR, "training_summary.json"),
"w",
encoding="utf-8",
) as file:
json.dump(metrics, file, indent=2)
output_volume.commit()
print("=" * 72)
print("TRAINING COMPLETE")
print(f"Final adapter: {final_adapter_dir}")
print(f"Elapsed time: {elapsed_seconds / 60:.2f} minutes")
print("=" * 72)
return {
"output_directory": OUTPUT_DIR,
"final_adapter": final_adapter_dir,
"elapsed_minutes": elapsed_seconds / 60,
}
@app.local_entrypoint()
def main():
result = train.remote()
print("\nRemote training finished.")
print(result)