AriaLM / src /s00_main.py
krishnah27's picture
Upload folder using huggingface_hub
30e9297 verified
Raw
History Blame Contribute Delete
5.62 kB
"""
Main entry point for Music Generation LLM.
Orchestrates: config → data → model → train → generate.
Usage:
python -m src.s00_main train # Train the model
python -m src.s00_main generate # Generate music from trained model
python -m src.s00_main train+generate # Train then generate
"""
import argparse
import logging
import sys
from pathlib import Path
import torch
# Add project root to path
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from src.s01_config import ModelConfig, TrainConfig, DataConfig, GenConfig, PathConfig, get_device
from src.s02_tokenizer import MusicTokenizer
from src.s03_dataset import create_dataloaders
from src.s04_model import MusicTransformer
from src.s05_trainer import Trainer
from src.s06_generator import generate_midi_file
from src.s07_utils import setup_logging, set_seed, log_memory_usage, clear_memory
logger = logging.getLogger(__name__)
def train_pipeline(
model_config: ModelConfig,
train_config: TrainConfig,
data_config: DataConfig,
path_config: PathConfig,
tokenizer: MusicTokenizer,
):
"""Full training pipeline."""
logger.info("=" * 60)
logger.info("MUSIC GENERATION LLM — TRAINING")
logger.info("=" * 60)
# Create data loaders
logger.info("Preparing data...")
train_loader, val_loader = create_dataloaders(
data_config, train_config, path_config, tokenizer
)
# Build model
model_config.vocab_size = tokenizer.vocab_size
model = MusicTransformer.from_config(model_config)
logger.info(f"Model: {model.count_parameters():,} parameters")
log_memory_usage("Pre-training")
# Train
trainer = Trainer(model, train_loader, val_loader, train_config, path_config)
trainer.train()
# Save tokenizer
tokenizer.save(path_config.tokenizer_path)
log_memory_usage("Post-training")
clear_memory()
return model
def generate_pipeline(
model_config: ModelConfig,
gen_config: GenConfig,
path_config: PathConfig,
tokenizer: MusicTokenizer,
model: MusicTransformer | None = None,
):
"""Generate music from trained model."""
logger.info("=" * 60)
logger.info("MUSIC GENERATION LLM — GENERATING")
logger.info("=" * 60)
if model is None:
best_ckpt = path_config.checkpoint_dir / "best.pt"
if not best_ckpt.exists():
logger.error(f"No checkpoint found at {best_ckpt}. Train first!")
return
model_config.vocab_size = tokenizer.vocab_size
model = MusicTransformer.from_config(model_config)
ckpt = torch.load(best_ckpt, map_location=get_device(), weights_only=False)
model.load_state_dict(ckpt["model_state_dict"])
logger.info(f"Loaded model from {best_ckpt}")
# Generate multiple samples
for i in range(3):
output_path = path_config.output_dir / f"generated_{i+1}.mid"
gen_config_i = GenConfig(
temperature=gen_config.temperature,
top_k=gen_config.top_k,
top_p=gen_config.top_p,
max_tokens=gen_config.max_tokens,
repetition_penalty=gen_config.repetition_penalty,
seed=gen_config.seed + i,
)
generate_midi_file(model, tokenizer, gen_config_i, output_path)
logger.info(f"Generated sample {i+1}: {output_path}")
clear_memory()
def main():
parser = argparse.ArgumentParser(description="Music Generation LLM")
parser.add_argument(
"mode",
choices=["train", "generate", "train+generate"],
default="train+generate",
nargs="?",
help="Operation mode",
)
parser.add_argument("--epochs", type=int, default=None, help="Override max epochs")
parser.add_argument("--batch-size", type=int, default=None, help="Override batch size")
parser.add_argument("--lr", type=float, default=None, help="Override learning rate")
parser.add_argument("--seq-len", type=int, default=None, help="Override max sequence length")
parser.add_argument("--temperature", type=float, default=None, help="Generation temperature")
parser.add_argument("--max-tokens", type=int, default=None, help="Max generation tokens")
args = parser.parse_args()
setup_logging()
set_seed(42)
# Initialize configs
model_config = ModelConfig()
train_config = TrainConfig()
data_config = DataConfig()
gen_config = GenConfig()
path_config = PathConfig()
# Apply overrides
if args.epochs:
train_config.max_epochs = args.epochs
if args.batch_size:
train_config.batch_size = args.batch_size
if args.lr:
train_config.learning_rate = args.lr
if args.seq_len:
data_config.max_seq_len = args.seq_len
model_config.max_seq_len = args.seq_len
if args.temperature:
gen_config.temperature = args.temperature
if args.max_tokens:
gen_config.max_tokens = args.max_tokens
# Tokenizer
tokenizer = MusicTokenizer()
model_config.vocab_size = tokenizer.vocab_size
logger.info(f"Device: {get_device()}")
logger.info(f"Vocab size: {tokenizer.vocab_size}")
logger.info(f"Model dim: {model_config.dim}, layers: {model_config.n_layers}, "
f"heads: {model_config.n_heads}, kv_heads: {model_config.n_kv_heads}")
model = None
if "train" in args.mode:
model = train_pipeline(model_config, train_config, data_config, path_config, tokenizer)
if "generate" in args.mode:
generate_pipeline(model_config, gen_config, path_config, tokenizer, model)
if __name__ == "__main__":
main()