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