CoolFace
Modelpublic

Spatial9/GravityLLM

sourceHugging Faceapache-2.0updated 7mo agoView on Hugging Face
0likes
train.py246 linesDownload Raw Back to root
1import argparse2import json3import os4import random5from dataclasses import dataclass6from pathlib import Path7from typing import Dict, List8 9import torch10from datasets import load_dataset11from peft import LoraConfig, get_peft_model, prepare_model_for_kbit_training12from transformers import (13    AutoModelForCausalLM,14    AutoTokenizer,15    BitsAndBytesConfig,16    Trainer,17    TrainingArguments,18    set_seed,19)20 21SYSTEM_PREFIX = (22    "You are GravityLLM, a Spatial9 scene generation model. "23    "Given music constraints and stem features, output ONLY valid Spatial9Scene JSON. "24    "Do not return markdown. Do not explain your answer. "25    "Respect hard constraints such as object budgets, anchor positions, and low-end centering.\n\n"26)27 28 29def parse_args() -> argparse.Namespace:30    parser = argparse.ArgumentParser(description="Fine-tune GravityLLM for Spatial9 scene generation.")31    parser.add_argument("--model", type=str, default="Qwen/Qwen2.5-1.5B-Instruct")32    parser.add_argument("--train_file", type=str, default="data/train.jsonl")33    parser.add_argument("--valid_file", type=str, default="data/valid.jsonl")34    parser.add_argument("--output_dir", type=str, default="outputs/GravityLLM-Qwen2.5-1.5B-S9")35    parser.add_argument("--max_length", type=int, default=2048)36 37    parser.add_argument("--num_train_epochs", type=float, default=1.0)38    parser.add_argument("--learning_rate", type=float, default=2e-4)39    parser.add_argument("--train_batch_size", type=int, default=1)40    parser.add_argument("--eval_batch_size", type=int, default=1)41    parser.add_argument("--gradient_accumulation_steps", type=int, default=16)42    parser.add_argument("--warmup_ratio", type=float, default=0.03)43    parser.add_argument("--weight_decay", type=float, default=0.0)44    parser.add_argument("--logging_steps", type=int, default=10)45    parser.add_argument("--save_steps", type=int, default=200)46    parser.add_argument("--eval_steps", type=int, default=200)47    parser.add_argument("--seed", type=int, default=42)48 49    parser.add_argument("--lora", action="store_true", help="Enable LoRA adapters.")50    parser.add_argument("--qlora", action="store_true", help="Enable 4-bit QLoRA training.")51    parser.add_argument("--lora_r", type=int, default=16)52    parser.add_argument("--lora_alpha", type=int, default=32)53    parser.add_argument("--lora_dropout", type=float, default=0.05)54 55    parser.add_argument("--bf16", action="store_true")56    parser.add_argument("--fp16", action="store_true")57 58    parser.add_argument("--push_to_hub", action="store_true")59    parser.add_argument("--hub_model_id", type=str, default=None)60    parser.add_argument("--hub_private_repo", action="store_true")61    return parser.parse_args()62 63 64def load_jsonl(file_path: str):65    return load_dataset("json", data_files=file_path, split="train")66 67 68def format_prompt(raw_prompt: str) -> str:69    raw_prompt = raw_prompt.strip()70    if raw_prompt.lower().startswith("gravityllm:"):71        raw_prompt = raw_prompt.split(":", 1)[1].strip()72    return SYSTEM_PREFIX + raw_prompt + "\n\nOUTPUT:\n"73 74 75def tokenize_example(example: Dict[str, str], tokenizer, max_length: int) -> Dict[str, List[int]]:76    prompt_text = format_prompt(example["prompt"])77    completion_text = example["completion"].strip()78 79    prompt_ids = tokenizer(prompt_text, add_special_tokens=False)["input_ids"]80    completion_ids = tokenizer(completion_text + tokenizer.eos_token, add_special_tokens=False)["input_ids"]81 82    input_ids = prompt_ids + completion_ids83    labels = [-100] * len(prompt_ids) + completion_ids84 85    if len(input_ids) > max_length:86        input_ids = input_ids[:max_length]87        labels = labels[:max_length]88 89    attention_mask = [1] * len(input_ids)90    return {91        "input_ids": input_ids,92        "attention_mask": attention_mask,93        "labels": labels,94    }95 96 97@dataclass98class CausalDataCollator:99    pad_token_id: int100    label_pad_token_id: int = -100101 102    def __call__(self, features):103        max_len = max(len(f["input_ids"]) for f in features)104 105        input_ids = []106        attention_mask = []107        labels = []108 109        for f in features:110            pad_len = max_len - len(f["input_ids"])111            input_ids.append(f["input_ids"] + [self.pad_token_id] * pad_len)112            attention_mask.append(f["attention_mask"] + [0] * pad_len)113            labels.append(f["labels"] + [self.label_pad_token_id] * pad_len)114 115        batch = {116            "input_ids": torch.tensor(input_ids, dtype=torch.long),117            "attention_mask": torch.tensor(attention_mask, dtype=torch.long),118            "labels": torch.tensor(labels, dtype=torch.long),119        }120        return batch121 122 123def prepare_model(args: argparse.Namespace):124    model_kwargs = {}125    if args.qlora:126        compute_dtype = torch.bfloat16 if args.bf16 else torch.float16127        model_kwargs["quantization_config"] = BitsAndBytesConfig(128            load_in_4bit=True,129            bnb_4bit_quant_type="nf4",130            bnb_4bit_use_double_quant=True,131            bnb_4bit_compute_dtype=compute_dtype,132        )133        model_kwargs["device_map"] = "auto"134 135    model = AutoModelForCausalLM.from_pretrained(136        args.model,137        torch_dtype=torch.bfloat16 if args.bf16 else (torch.float16 if args.fp16 else None),138        trust_remote_code=True,139        **model_kwargs,140    )141    model.config.use_cache = False142 143    if args.qlora:144        model = prepare_model_for_kbit_training(model)145 146    if args.lora or args.qlora:147        lora_config = LoraConfig(148            r=args.lora_r,149            lora_alpha=args.lora_alpha,150            lora_dropout=args.lora_dropout,151            bias="none",152            task_type="CAUSAL_LM",153            target_modules="all-linear",154        )155        model = get_peft_model(model, lora_config)156        model.print_trainable_parameters()157 158    return model159 160 161def main() -> None:162    args = parse_args()163    os.makedirs(args.output_dir, exist_ok=True)164    set_seed(args.seed)165 166    tokenizer = AutoTokenizer.from_pretrained(args.model, use_fast=True, trust_remote_code=True)167    tokenizer.padding_side = "right"168    if tokenizer.pad_token is None:169        tokenizer.pad_token = tokenizer.eos_token170 171    train_ds = load_jsonl(args.train_file)172    valid_ds = load_jsonl(args.valid_file) if args.valid_file and Path(args.valid_file).exists() else None173 174    train_ds = train_ds.map(175        lambda row: tokenize_example(row, tokenizer, args.max_length),176        remove_columns=train_ds.column_names,177        desc="Tokenizing train set",178    )179    if valid_ds is not None:180        valid_ds = valid_ds.map(181            lambda row: tokenize_example(row, tokenizer, args.max_length),182            remove_columns=valid_ds.column_names,183            desc="Tokenizing valid set",184        )185 186    model = prepare_model(args)187 188    training_args = TrainingArguments(189        output_dir=args.output_dir,190        overwrite_output_dir=True,191        num_train_epochs=args.num_train_epochs,192        learning_rate=args.learning_rate,193        per_device_train_batch_size=args.train_batch_size,194        per_device_eval_batch_size=args.eval_batch_size,195        gradient_accumulation_steps=args.gradient_accumulation_steps,196        warmup_ratio=args.warmup_ratio,197        weight_decay=args.weight_decay,198        logging_steps=args.logging_steps,199        save_steps=args.save_steps,200        eval_steps=args.eval_steps,201        evaluation_strategy="steps" if valid_ds is not None else "no",202        save_strategy="steps",203        bf16=args.bf16,204        fp16=args.fp16,205        report_to="none",206        gradient_checkpointing=True,207        lr_scheduler_type="cosine",208        optim="paged_adamw_32bit" if (args.lora or args.qlora) else "adamw_torch",209        max_grad_norm=1.0,210        push_to_hub=args.push_to_hub,211        hub_model_id=args.hub_model_id,212        hub_private_repo=args.hub_private_repo,213        hub_strategy="end" if args.push_to_hub else "every_save",214    )215 216    trainer = Trainer(217        model=model,218        args=training_args,219        train_dataset=train_ds,220        eval_dataset=valid_ds,221        data_collator=CausalDataCollator(pad_token_id=tokenizer.pad_token_id),222        tokenizer=tokenizer,223    )224 225    train_result = trainer.train()226    trainer.save_model(args.output_dir)227    tokenizer.save_pretrained(args.output_dir)228 229    metrics = train_result.metrics230    with open(Path(args.output_dir) / "training_metrics.json", "w", encoding="utf-8") as f:231        json.dump(metrics, f, indent=2)232 233    run_meta = vars(args).copy()234    run_meta["train_examples"] = len(train_ds)235    run_meta["valid_examples"] = len(valid_ds) if valid_ds is not None else 0236    with open(Path(args.output_dir) / "run_config.json", "w", encoding="utf-8") as f:237        json.dump(run_meta, f, indent=2)238 239    if args.push_to_hub:240        trainer.push_to_hub(commit_message="Add GravityLLM fine-tuned adapter")241    print(f"Training complete. Artifacts saved to: {args.output_dir}")242 243 244if __name__ == "__main__":245    main()246