add HF inference script, enrich readme w/ egs

This commit is contained in:
slSeanWU
2025-10-30 17:02:58 -04:00
parent d942658721
commit 782ec8cbac
2 changed files with 551 additions and 18 deletions
+76 -17
View File
@@ -4,7 +4,7 @@ Built on **Llama 3.2 (1B)** with an extended vocabulary for MIDI tokens.
## Research Paper ## Research Paper
- Shih-Lun Wu, Yoon Kim, and Cheng-Zhi Anna Huang. - Shih-Lun Wu, Yoon Kim, and Cheng-Zhi Anna Huang.
"**MIDI-LLM: Adapting large language models for text-to-MIDI music generation**." "**MIDI-LLM: Adapting Large Language Models for Text-to-MIDI Music Generation**."
NeurIPS AI4Music Workshop, 2025. NeurIPS AI4Music Workshop, 2025.
## Setup ## Setup
@@ -54,46 +54,54 @@ pip install -r requirements.txt
python -c "import torch; from vllm import LLM; from anticipation.convert import events_to_midi; print('Setup successful')" python -c "import torch; from vllm import LLM; from anticipation.convert import events_to_midi; print('Setup successful')"
``` ```
## Run Inference with vLLM ## Inference (Generation) Usage
### Example 1: Single prompt
**IMPORTANT**: We provide two inference backends with different trade-offs:
- **vLLM** (`generate_vllm.py`): Faster token generation but more complex setup and longer initialization. **Recommended for batch inference (multiple prompts) or interactive sessions.**
- **Transformers** (`generate_transformers.py`): Simpler setup and faster initialization, but slower generation. **Recommended for quick single-prompt testing.**
Both scripts share the same arguments (except for `--fp8` quantization, which only works in vLLM) and output format.
### Example 1: Single prompt (use transformers)
```bash ```bash
python generate_vllm.py \ python generate_transformers.py \
--model slseanwu/MIDI-LLM_Llama-3.2-1B # will pull from huggingface hub \
--prompt "A cheerful piano melody" --prompt "A cheerful piano melody"
``` ```
This will output 4 MIDIs (and the synthesized MP3s) conditioned on the same input prompt Outputs 4 MIDIs (and synthesized MP3s) conditioned on the same prompt by default.
### Example 2: Batch generation from file ### Example 2: Batch generation from file (use vLLM)
```bash ```bash
python generate_vllm.py \ python generate_vllm.py \
--model slseanwu/MIDI-LLM_Llama-3.2-1B \
--prompts_file some_example_prompts.txt \ --prompts_file some_example_prompts.txt \
--fp8 \ --fp8 \
--no-synthesize --no-synthesize
``` ```
- `some_example_prompts.txt` should contain one prompt per line. - `some_example_prompts.txt` should contain one prompt per line.
- `--fp8` performs dynamic weight quantization for faster inference. - `--fp8` performs FP8 quantization for faster inference.
- `--no-synthesize` skips audio synthesis (i.e., outputs MIDI only). - `--no-synthesize` skips audio synthesis (outputs MIDI only).
### Example 3: Interactive mode ### Example 3: Interactive mode (use vLLM)
```bash ```bash
python generate_vllm.py \ python generate_vllm.py \
--model slseanwu/MIDI-LLM_Llama-3.2-1B \ --interactive \
--output_root generations_interactive/ \ --output_root generations_interactive/ \
--interactive --n_outputs 1
``` ```
- Outputs will be saved under `generations_interactive/` Loads the model once, then lets you enter prompts continuously. Press Enter with an empty prompt to exit.
This loads the model once, then lets you enter prompts interactively. Press Enter with empty prompt to exit.
- Outputs will be stored under `generations_interactive/`
- `--n_outputs 1` generates only 1 output for each prompt
### More options ### More options
See full options with: See full options for either script with:
```bash ```bash
python generate_transformers.py --help # or
python generate_vllm.py --help python generate_vllm.py --help
``` ```
### Inference Output Structure ### Inference output structure
``` ```
[output_root]/ [output_root]/
└── 2025-10-30_143022/ # Session timestamp └── 2025-10-30_143022/ # Session timestamp
@@ -104,3 +112,54 @@ python generate_vllm.py --help
│ └── ... │ └── ...
└── generation_stats.json └── generation_stats.json
``` ```
## Example Prompts
Here are some example prompts to get you started. The model can work with both detailed descriptions similar to what's seen at training, and creative free-form prompts.
### In-Domain Examples (from validation set)
<details>
<summary><b>Example 1: Rock with pop influence</b></summary>
```
A melodic and energetic rock song with a touch of pop influence, featuring synth
strings, piano, distortion guitar, synth voice, and drums, all contributing to a
blend of happy and dark moods. Set in the key of A minor with a 4/4 time signature,
this fast-paced track showcases a chord progression of Bm, Cmaj7, and Gmaj7.
```
</details>
<details>
<summary><b>Example 2: Classical soundtrack</b></summary>
```
A slow and relaxing classical piece featuring a church organ and French horn, likely
to be used as a soundtrack in a dramatic or emotional film. Written in A minor and 4/4
time. The chord progression of E7, Am, and E contributes to the piece's sentimental
atmosphere.
```
</details>
### Creative Custom Prompts
<details>
<summary><b>Example 3: Road trip song</b></summary>
```
An energetic and motivating pop song you love to hear on a long road trip.
```
</details>
<details>
<summary><b>Example 4: Sunday picnic jazz</b></summary>
```
Upbeat and playful jazz music with lively saxophones, like you're going out on a
Sunday picnic.
```
</details>
+474
View File
@@ -0,0 +1,474 @@
#!/usr/bin/env python3
"""
MIDI-LLM: Text-to-MIDI Generation using HuggingFace Transformers
This script generates MIDI files from text prompts using the MIDI-LLM model with HuggingFace backend.
Simpler to set up than vLLM but slower for inference.
"""
import json
import time
import argparse
from pathlib import Path
from datetime import datetime
from typing import List, Optional
import torch
import tqdm
from transformers import AutoTokenizer, AutoModelForCausalLM
# Import helper functions and constants
from midi_llm.utils import (
save_generation,
AMT_GPT2_BOS_ID,
LLAMA_VOCAB_SIZE,
LLAMA_MODEL_NAME,
)
# Default generation parameters
DEFAULT_TEMPERATURE = 1.0
DEFAULT_TOP_P = 0.98
DEFAULT_MAX_TOKENS = 2046
DEFAULT_N_OUTPUTS = 4 # give more outputs for variability
def prepare_hf_model(model_path: str):
"""
Initialize HuggingFace model in BFloat16.
Args:
model_path: Path to model checkpoint
Returns:
model
"""
print(f"\n{'='*70}")
print("Model Configuration")
print(f"{'='*70}")
print(f"Model path: {model_path}")
print(f"Precision: BFloat16")
print(f"{'='*70}\n")
# Load model in BF16
model = AutoModelForCausalLM.from_pretrained(
model_path,
dtype=torch.bfloat16,
trust_remote_code=True
).to(device="cuda")
model.eval()
print(f"✓ Model loaded successfully\n")
return model
def generate_from_prompts_hf(
model,
tokenizer: AutoTokenizer,
prompts: List[str],
output_dir: Path,
model_path: str,
soundfont_path: Optional[str] = None,
synthesize: bool = False,
temperature: float = DEFAULT_TEMPERATURE,
top_p: float = DEFAULT_TOP_P,
max_tokens: int = DEFAULT_MAX_TOKENS,
n_outputs: int = DEFAULT_N_OUTPUTS,
system_prompt: Optional[str] = None
) -> dict:
"""
Generate MIDI from text prompts using HuggingFace model.
Args:
model: HuggingFace model
tokenizer: HuggingFace tokenizer
prompts: List of text prompts
output_dir: Base output directory
model_path: Path to model (to check for with_edits)
soundfont_path: Path to SoundFont file
synthesize: Whether to synthesize to audio
temperature: Sampling temperature
top_p: Nucleus sampling threshold
max_tokens: Maximum tokens to generate
n_outputs: Number of outputs per prompt
system_prompt: Optional system prompt prefix
Returns:
Dictionary with generation statistics and output files
"""
# Default system prompt
if system_prompt is None:
system_prompt = "You are a world-class composer. Please compose some music according to the following description: "
stats = {
"total_prompts": len(prompts),
"successful_generations": 0,
"failed_generations": 0,
"generation_times": [],
"output_files": []
}
for idx, prompt in enumerate(tqdm.tqdm(prompts, desc="Generating")):
print(f"\n[{idx+1}/{len(prompts)}] Prompt: {prompt}")
# Create output directory for this prompt
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
prompt_output_dir = output_dir / f"{timestamp}_prompt_{idx+1}"
# Prepare full prompt, add space to the end of each prompt to match training
full_prompt = system_prompt + prompt + " "
# Tokenize
llama_input = tokenizer(full_prompt, return_tensors="pt", padding=False)
input_ids = llama_input["input_ids"]
# Add MIDI BOS token
midi_bos = torch.tensor([[AMT_GPT2_BOS_ID + LLAMA_VOCAB_SIZE]])
input_ids = torch.cat([input_ids, midi_bos], dim=1)
# Move to device
device = next(model.parameters()).device
input_ids = input_ids.to(device)
# Generate multiple outputs
start_time = time.time()
with torch.no_grad():
outputs = model.generate(
input_ids=input_ids,
do_sample=True,
max_new_tokens=max_tokens,
temperature=temperature,
top_p=top_p,
num_return_sequences=n_outputs,
pad_token_id=tokenizer.pad_token_id,
)
generation_time = time.time() - start_time
if idx > 0: # Skip first generation for timing (warmup)
stats["generation_times"].append(generation_time)
print(f"Generation time: {generation_time:.2f}s")
# Extract only the generated tokens (remove prompt)
prompt_len = input_ids.shape[1]
outputs = outputs[:, prompt_len:]
# Shift tokens back to MIDI vocab range
outputs = outputs - LLAMA_VOCAB_SIZE
outputs = outputs.cpu().tolist()
# Save all outputs for this prompt
successful_outputs = 0
prompt_files = []
for output_idx, midi_tokens in enumerate(outputs):
# Save generation
success = save_generation(
tokens=midi_tokens,
prompt=prompt,
output_dir=prompt_output_dir,
generation_idx=output_idx + 1,
soundfont_path=soundfont_path,
synthesize=synthesize
)
if success:
successful_outputs += 1
# Track output files
midi_file = prompt_output_dir / f"gen_{output_idx + 1}.mid"
prompt_files.append(str(midi_file))
if synthesize and soundfont_path:
mp3_file = prompt_output_dir / f"gen_{output_idx + 1}.mp3"
if mp3_file.exists():
prompt_files.append(str(mp3_file))
print(f"Successfully saved {successful_outputs}/{n_outputs} outputs")
stats["successful_generations"] += successful_outputs
stats["failed_generations"] += (n_outputs - successful_outputs)
stats["output_files"].extend(prompt_files)
return stats
def main():
parser = argparse.ArgumentParser(
description="Generate MIDI files from text prompts using MIDI-LLM with HuggingFace",
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog="""
Examples:
# Generate from a single prompt (there will be 4 outputs by default)
python generate_transformers.py --model path/to/checkpoint \\
--prompt "A cheerful piano melody"
# Generate single output without synthesis
python generate_transformers.py --model path/to/checkpoint \\
--prompt "A relaxing jazz piece" \\
--n_outputs 1 \\
--no-synthesize
# Interactive mode (with initial prompt)
python generate_transformers.py --model path/to/checkpoint \\
--prompt "A cheerful melody" \\
--interactive
# Interactive-only mode (no initial prompt)
python generate_transformers.py --model path/to/checkpoint \\
--interactive
# Generate from prompts file
python generate_transformers.py --model path/to/checkpoint \\
--prompts_file prompts.txt
"""
)
# Required arguments
parser.add_argument(
"--model",
type=str,
default="slseanwu/MIDI-LLM_Llama-3.2-1B",
help="Path to MIDI-LLM model checkpoint, can be HuggingFace model ID or local path (default: slseanwu/MIDI-LLM_Llama-3.2-1B)"
)
# Input arguments (not required if using --interactive only)
input_group = parser.add_mutually_exclusive_group(required=False)
input_group.add_argument(
"--prompt",
type=str,
help="Single text prompt for generation"
)
input_group.add_argument(
"--prompts_file",
type=str,
help="Path to file containing prompts (one per line)"
)
# Output arguments
parser.add_argument(
"--output_root",
type=str,
default="./generated_outputs",
help="Root directory for outputs (timestamped subdirs will be created inside, default: ./generated_outputs)"
)
parser.add_argument(
"--n_outputs",
type=int,
default=DEFAULT_N_OUTPUTS,
help=f"Number of outputs to generate per prompt (default: {DEFAULT_N_OUTPUTS})"
)
# Synthesis arguments
parser.add_argument(
"--no-synthesize",
dest="synthesize",
action="store_false",
help="Skip audio synthesis (only generate MIDI files)"
)
parser.set_defaults(synthesize=True)
parser.add_argument(
"--soundfont",
type=str,
default="./soundfonts/FluidR3_GM/FluidR3_GM.sf2",
help="Path to SoundFont file for synthesis (default: ./soundfonts/FluidR3_GM/FluidR3_GM.sf2)"
)
# Generation parameters
parser.add_argument(
"--temperature",
type=float,
default=DEFAULT_TEMPERATURE,
help=f"Sampling temperature (default: {DEFAULT_TEMPERATURE})"
)
parser.add_argument(
"--top_p",
type=float,
default=DEFAULT_TOP_P,
help=f"Nucleus sampling threshold (default: {DEFAULT_TOP_P})"
)
parser.add_argument(
"--max_tokens",
type=int,
default=DEFAULT_MAX_TOKENS,
help=f"Maximum tokens to generate (default: {DEFAULT_MAX_TOKENS})"
)
# Model arguments
parser.add_argument(
"--cache_dir",
type=str,
default=None,
help="HuggingFace cache directory (default: $HF_HOME or ~/.cache/huggingface)"
)
parser.add_argument(
"--interactive",
action="store_true",
help="Enter interactive mode after initial generation (keep generating until empty prompt)"
)
args = parser.parse_args()
# Validate that either prompts are provided or interactive mode is enabled
if not args.prompt and not args.prompts_file and not args.interactive:
parser.error("Either --prompt, --prompts_file, or --interactive must be specified")
# Load prompts (if provided)
prompts = []
if args.prompt:
prompts = [args.prompt]
elif args.prompts_file:
with open(args.prompts_file, "r") as f:
prompts = [line.strip() for line in f if line.strip()]
print(f"Loaded {len(prompts)} prompts from {args.prompts_file}")
# Check synthesis requirements
if args.synthesize:
soundfont_path = Path(args.soundfont)
if not soundfont_path.exists():
print(f"Error: SoundFont not found at {soundfont_path}")
print("Please download a SoundFont or disable synthesis")
import sys
sys.exit(1)
# Create output root directory with timestamp
output_root = Path(args.output_root)
session_timestamp = datetime.now().strftime("%Y-%m-%d_%H%M%S")
output_dir = output_root / session_timestamp
output_dir.mkdir(parents=True, exist_ok=True)
print(f"Output directory: {output_dir.absolute()}\n")
# Load tokenizer from the model checkpoint
print("Loading tokenizer...")
tokenizer = AutoTokenizer.from_pretrained(
args.model,
cache_dir=args.cache_dir,
pad_token="<|eot_id|>",
)
# Load model
model = prepare_hf_model(model_path=args.model)
# Generate from initial prompts (if provided)
if prompts:
print(f"Starting generation for {len(prompts)} prompt(s)...\n")
start_time = time.time()
stats = generate_from_prompts_hf(
model=model,
tokenizer=tokenizer,
prompts=prompts,
output_dir=output_dir,
model_path=args.model,
soundfont_path=args.soundfont if args.synthesize else None,
synthesize=args.synthesize,
temperature=args.temperature,
top_p=args.top_p,
max_tokens=args.max_tokens,
n_outputs=args.n_outputs
)
total_time = time.time() - start_time
# Print summary
print(f"\n{'='*70}")
print("Generation Summary")
print(f"{'='*70}")
print(f"Total prompts: {stats['total_prompts']}")
print(f"Successful generations: {stats['successful_generations']}")
print(f"Failed generations: {stats['failed_generations']}")
print(f"Total time: {total_time:.2f}s")
if stats['generation_times']:
avg_time = sum(stats['generation_times']) / len(stats['generation_times'])
print(f"Average generation time: {avg_time:.2f}s (excluding warmup)")
print(f"\nOutputs saved to: {output_dir.absolute()}")
# Print generated files
if stats['output_files']:
print(f"\nGenerated files:")
for file_path in stats['output_files']:
file_type = "🎵 MIDI" if file_path.endswith('.mid') else "🎧 Audio"
print(f" {file_type}: {file_path}")
print(f"{'='*70}\n")
# Save stats to JSON
stats_file = output_dir / "generation_stats.json"
with open(stats_file, "w") as f:
json.dump({
**stats,
"total_time": total_time,
"average_time": sum(stats['generation_times']) / len(stats['generation_times']) if stats['generation_times'] else 0,
"config": {
"model": args.model,
"temperature": args.temperature,
"top_p": args.top_p,
"max_tokens": args.max_tokens,
"n_outputs": args.n_outputs,
}
}, f, indent=2)
else:
print(f"No initial prompts provided. Starting in interactive mode...\n")
# Interactive mode
if args.interactive:
print(f"\n{'='*70}")
print("Interactive Mode")
print(f"{'='*70}")
print("Enter prompts to generate more MIDI files.")
print("Press Enter with empty prompt to exit.\n")
while True:
try:
# Get user input
user_prompt = input("Prompt: ").strip()
# Exit if empty
if not user_prompt:
print("\nExiting interactive mode. Goodbye!")
break
# Generate from the new prompt
print()
interactive_stats = generate_from_prompts_hf(
model=model,
tokenizer=tokenizer,
prompts=[user_prompt],
output_dir=output_dir,
model_path=args.model,
soundfont_path=args.soundfont if args.synthesize else None,
synthesize=args.synthesize,
temperature=args.temperature,
top_p=args.top_p,
max_tokens=args.max_tokens,
n_outputs=args.n_outputs
)
# Print input prompt
print(f"Input prompt: {user_prompt}")
# Print mini summary
print(f"\n✓ Generated {interactive_stats['successful_generations']}/{args.n_outputs} outputs")
if interactive_stats['generation_times']:
print(f" Generation time: {interactive_stats['generation_times'][0]:.2f}s")
# Print file paths
if interactive_stats['output_files']:
for file_path in interactive_stats['output_files']:
file_type = "🎵" if file_path.endswith('.mid') else "🎧"
print(f" {file_type} {file_path}")
print()
except KeyboardInterrupt:
print("\n\nInterrupted. Exiting interactive mode.")
break
except EOFError:
print("\n\nExiting interactive mode.")
break
if __name__ == "__main__":
main()