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()