| """ |
| 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 |
|
|
| |
| 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) |
|
|
| |
| logger.info("Preparing data...") |
| train_loader, val_loader = create_dataloaders( |
| data_config, train_config, path_config, tokenizer |
| ) |
|
|
| |
| 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") |
|
|
| |
| trainer = Trainer(model, train_loader, val_loader, train_config, path_config) |
| trainer.train() |
|
|
| |
| 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}") |
|
|
| |
| 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) |
|
|
| |
| model_config = ModelConfig() |
| train_config = TrainConfig() |
| data_config = DataConfig() |
| gen_config = GenConfig() |
| path_config = PathConfig() |
|
|
| |
| 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 = 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() |
|
|