Raw History Blame Contribute Delete
10 kB
from __future__ import annotations
import argparse
from pathlib import Path
import torch
from accelerate import Accelerator, DistributedType
from datasets import load_dataset
from torch.distributed.tensor import DTensor
from torch.utils.data import DataLoader
from transformers import AutoConfig, AutoModelForCausalLM, AutoTokenizer
from peft import LoraConfig, TaskType, get_peft_model, get_peft_model_state_dict
DEFAULT_MODEL = Path(__file__).resolve().parent.parent
DATASET_NAME = "tatsu-lab/alpaca"
DATASET_REVISION = "dce01c9b08f87459cf36a430d809084718273017"
FSDP_VERSION = 2
LORA_TARGET_MODULES = ["q_proj", "k_proj", "v_proj", "o_proj"]
MAX_RESPONSE_TOKENS = 128
ROUTER_BUFFER_NAME = "e_score_correction_bias"
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="LoRA fine-tuning for Alice AI")
parser.add_argument("--model", default=str(DEFAULT_MODEL))
parser.add_argument("--output-dir", type=Path)
parser.add_argument("--steps", type=int, default=1)
parser.add_argument("--sequence-length", type=int, default=32)
parser.add_argument("--learning-rate", type=float, default=2e-4)
parser.add_argument("--lora-rank", type=int, default=8)
parser.add_argument("--seed", type=int, default=42)
return parser.parse_args()
def encode_alpaca_row(row, tokenizer, sequence_length: int) -> dict[str, list[int]]:
prompt = f"Instruction:\n{row['instruction']}"
if row["input"]:
prompt += f"\n\nInput:\n{row['input']}"
prompt += "\n\nResponse:\n"
prompt_ids = tokenizer(prompt, add_special_tokens=False).input_ids
response_ids = tokenizer(row["output"], add_special_tokens=False).input_ids
response_budget = min(MAX_RESPONSE_TOKENS, max(1, sequence_length // 4))
response_ids = response_ids[:response_budget]
prompt_ids = prompt_ids[: sequence_length - len(response_ids) - 2]
input_ids = [
tokenizer.bos_token_id,
*prompt_ids,
*response_ids,
tokenizer.eos_token_id,
]
labels = [-100] * (len(prompt_ids) + 1) + [*response_ids, tokenizer.eos_token_id]
attention_mask = [1] * len(input_ids)
padding_length = sequence_length - len(input_ids)
input_ids.extend([tokenizer.pad_token_id] * padding_length)
labels.extend([-100] * padding_length)
attention_mask.extend([0] * padding_length)
return {
"input_ids": input_ids,
"attention_mask": attention_mask,
"labels": labels,
}
def build_dataloader(tokenizer, sequence_length: int) -> DataLoader:
if tokenizer.pad_token_id is None:
tokenizer.pad_token = tokenizer.eos_token
dataset = load_dataset(
DATASET_NAME,
revision=DATASET_REVISION,
split="train",
)
dataset = dataset.map(
encode_alpaca_row,
fn_kwargs={"tokenizer": tokenizer, "sequence_length": sequence_length},
remove_columns=dataset.column_names,
)
return DataLoader(dataset.with_format("torch"), batch_size=1, shuffle=True)
def save_adapter(accelerator: Accelerator, model, tokenizer, output_dir: Path) -> None:
unwrapped_model = accelerator.unwrap_model(model)
adapter_state = get_peft_model_state_dict(unwrapped_model)
adapter_state = {
name: value.full_tensor().cpu() if isinstance(value, DTensor) else value.cpu()
for name, value in adapter_state.items()
}
if accelerator.is_main_process:
unwrapped_model.save_pretrained(output_dir, state_dict=adapter_state)
tokenizer.save_pretrained(output_dir)
accelerator.wait_for_everyone()
def restore_rotary_buffer(
accelerator: Accelerator,
model,
device: torch.device | None = None,
) -> None:
unwrapped_model = accelerator.unwrap_model(model)
config = unwrapped_model.config
base_model = (
unwrapped_model.get_base_model()
if hasattr(unwrapped_model, "get_base_model")
else unwrapped_model
)
rotary_embedding = base_model.model.rotary_emb
rotary_dim = int(config.head_dim * config.partial_rotary_factor)
inv_freq = 1.0 / (
config.rope_theta
** (
torch.arange(
0,
rotary_dim,
2,
dtype=torch.float32,
device=device or rotary_embedding.inv_freq.device,
)
/ rotary_dim
)
)
rotary_embedding.inv_freq = inv_freq
def load_model(model_path: str, accelerator: Accelerator):
load_kwargs = {
"trust_remote_code": True,
"dtype": torch.bfloat16,
"attn_implementation": "flash_attention_2",
}
if (
not accelerator.state.fsdp_plugin.cpu_ram_efficient_loading
or accelerator.is_main_process
):
return AutoModelForCausalLM.from_pretrained(model_path, **load_kwargs)
config = AutoConfig.from_pretrained(model_path, trust_remote_code=True)
# Transformers validates FlashAttention against a real device even though
# the attention backend does not affect the module layout.
config._attn_implementation = "eager"
previous_dtype = torch.get_default_dtype()
torch.set_default_dtype(load_kwargs["dtype"])
try:
with torch.device("meta"):
model = AutoModelForCausalLM.from_config(config, trust_remote_code=True)
finally:
torch.set_default_dtype(previous_dtype)
model.config._attn_implementation = load_kwargs["attn_implementation"]
return model
def remove_router_buffers(model, is_main_process: bool) -> dict[str, torch.Tensor]:
router_buffers = {}
for module_name, module in model.named_modules():
buffer = module._buffers.pop(ROUTER_BUFFER_NAME, None)
if buffer is None:
continue
router_buffers[module_name] = (
buffer.detach().cpu()
if is_main_process
else torch.empty(buffer.shape, dtype=buffer.dtype)
)
return router_buffers
def restore_router_buffers(
accelerator: Accelerator,
model,
router_buffers: dict[str, torch.Tensor],
) -> None:
unwrapped_model = accelerator.unwrap_model(model)
for module_name, buffer in router_buffers.items():
device_buffer = buffer.to(accelerator.device)
torch.distributed.broadcast(device_buffer, src=0)
unwrapped_model.get_submodule(module_name).register_buffer(
ROUTER_BUFFER_NAME,
device_buffer,
)
def materialize_trainable_parameters(model) -> None:
# Accelerate 1.14 remaps FSDP2 optimizer parameters by data_ptr(). All meta
# tensors use pointer 0, so give the small trainable LoRA tensors real storage.
for module in model.modules():
for parameter_name, parameter in tuple(module.named_parameters(recurse=False)):
if parameter.requires_grad and parameter.is_meta:
module._parameters[parameter_name] = torch.nn.Parameter(
torch.empty_like(parameter, device="cpu"),
)
def main() -> None:
args = parse_args()
if args.steps < 1:
raise ValueError("--steps must be positive")
if args.sequence_length <= 1:
raise ValueError("--sequence-length must be at least 2")
accelerator = Accelerator()
if accelerator.distributed_type != DistributedType.FSDP:
raise RuntimeError("Launch this script with Accelerate FSDP")
if accelerator.state.fsdp_plugin.fsdp_version != FSDP_VERSION:
raise RuntimeError("This example requires FSDP2")
torch.manual_seed(args.seed)
tokenizer = AutoTokenizer.from_pretrained(args.model, trust_remote_code=True)
with accelerator.main_process_first():
dataloader = build_dataloader(tokenizer, args.sequence_length)
lora_config = LoraConfig(
task_type=TaskType.CAUSAL_LM,
r=args.lora_rank,
lora_alpha=2 * args.lora_rank,
lora_dropout=0.0,
target_modules=LORA_TARGET_MODULES,
bias="none",
)
model = load_model(args.model, accelerator)
if accelerator.state.fsdp_plugin.cpu_ram_efficient_loading:
restore_rotary_buffer(accelerator, model, device=torch.device("cpu"))
model.config.use_cache = False
model = get_peft_model(model, lora_config)
if accelerator.state.fsdp_plugin.cpu_ram_efficient_loading:
materialize_trainable_parameters(model)
router_buffers = (
remove_router_buffers(model, accelerator.is_main_process)
if accelerator.state.fsdp_plugin.cpu_ram_efficient_loading
else {}
)
model.gradient_checkpointing_enable(
gradient_checkpointing_kwargs={"use_reentrant": False}
)
trainable = sum(
parameter.numel() for parameter in model.parameters() if parameter.requires_grad
)
total = sum(parameter.numel() for parameter in model.parameters())
accelerator.print(f"Trainable parameters: {trainable:,} / {total:,}")
optimizer = torch.optim.AdamW(
(parameter for parameter in model.parameters() if parameter.requires_grad),
lr=args.learning_rate,
)
model, optimizer, dataloader = accelerator.prepare(model, optimizer, dataloader)
restore_router_buffers(accelerator, model, router_buffers)
restore_rotary_buffer(accelerator, model)
model.train()
data_iterator = iter(dataloader)
for step in range(args.steps):
try:
batch = next(data_iterator)
except StopIteration:
data_iterator = iter(dataloader)
batch = next(data_iterator)
optimizer.zero_grad(set_to_none=True)
outputs = model(**batch, use_cache=False)
accelerator.backward(outputs.loss)
optimizer.step()
accelerator.print(
f"step={step + 1} loss={outputs.loss.detach().float().item():.6f}"
)
if args.output_dir is not None:
save_adapter(accelerator, model, tokenizer, args.output_dir)
accelerator.print("LoRA fine-tuning smoke test completed")
accelerator.end_training()
if __name__ == "__main__":
main()