Spatial9/GravityLLM
0
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 