AriaLM / src /s06_generator.py
krishnah27's picture
Upload folder using huggingface_hub
30e9297 verified
Raw
History Blame Contribute Delete
4.59 kB
"""
Music generation / inference with:
- Top-k / Top-p (nucleus) sampling
- Temperature scaling
- Repetition penalty
- KV-cache for efficient autoregressive generation
"""
import logging
from pathlib import Path
import torch
import torch.nn.functional as F
from src.s01_config import GenConfig, get_device
from src.s02_tokenizer import MusicTokenizer, BOS_TOKEN, EOS_TOKEN
from src.s04_model import MusicTransformer
logger = logging.getLogger(__name__)
def top_k_top_p_filter(logits: torch.Tensor, top_k: int, top_p: float) -> torch.Tensor:
"""Filter logits using top-k and nucleus (top-p) sampling."""
if top_k > 0:
top_k = min(top_k, logits.size(-1))
indices_to_remove = logits < torch.topk(logits, top_k)[0][..., -1, None]
logits[indices_to_remove] = float("-inf")
if top_p < 1.0:
sorted_logits, sorted_indices = torch.sort(logits, descending=True)
cumulative_probs = torch.cumsum(F.softmax(sorted_logits, dim=-1), dim=-1)
sorted_indices_to_remove = cumulative_probs > top_p
sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone()
sorted_indices_to_remove[..., 0] = 0
indices_to_remove = sorted_indices_to_remove.scatter(
dim=-1, index=sorted_indices, src=sorted_indices_to_remove
)
logits[indices_to_remove] = float("-inf")
return logits
def apply_repetition_penalty(logits: torch.Tensor, past_tokens: list[int], penalty: float):
"""Penalize tokens that have appeared recently."""
if penalty == 1.0 or not past_tokens:
return logits
# Only penalize last 64 tokens to avoid over-suppression
recent = past_tokens[-64:]
unique_tokens = set(recent)
for token_id in unique_tokens:
if logits[0, token_id] > 0:
logits[0, token_id] /= penalty
else:
logits[0, token_id] *= penalty
return logits
@torch.no_grad()
def generate(
model: MusicTransformer,
tokenizer: MusicTokenizer,
config: GenConfig,
prompt_tokens: list[int] | None = None,
device: torch.device | None = None,
) -> list[int]:
"""
Generate music tokens autoregressively using KV-cache.
Args:
model: Trained MusicTransformer
tokenizer: MusicTokenizer instance
config: Generation config (temperature, top_k, etc.)
prompt_tokens: Optional seed tokens (if None, starts with BOS)
device: Target device
Returns:
List of generated token IDs
"""
if device is None:
device = get_device()
model.eval()
model.reset_caches()
model.to(device)
if prompt_tokens is None:
prompt_tokens = [BOS_TOKEN]
generated = list(prompt_tokens)
# Set seed for reproducibility
if config.seed is not None:
torch.manual_seed(config.seed)
# Process prompt in one shot (prefill)
input_tensor = torch.tensor([generated], dtype=torch.long, device=device)
logits, _ = model(input_tensor, use_cache=True)
for step in range(config.max_tokens - len(generated)):
# Get logits for last position
next_logits = logits[:, -1, :] / max(config.temperature, 1e-8)
# Apply repetition penalty
next_logits = apply_repetition_penalty(
next_logits, generated, config.repetition_penalty
)
# Top-k / Top-p filtering
next_logits = top_k_top_p_filter(next_logits, config.top_k, config.top_p)
# Sample
probs = F.softmax(next_logits, dim=-1)
next_token = torch.multinomial(probs, num_samples=1).item()
generated.append(next_token)
# Stop on EOS
if next_token == EOS_TOKEN:
break
# Next step: single token with KV-cache
input_tensor = torch.tensor([[next_token]], dtype=torch.long, device=device)
logits, _ = model(input_tensor, use_cache=True)
model.reset_caches()
return generated
def generate_midi_file(
model: MusicTransformer,
tokenizer: MusicTokenizer,
config: GenConfig,
output_path: Path,
prompt_tokens: list[int] | None = None,
):
"""Generate a MIDI file and save to disk."""
logger.info(f"Generating music (max_tokens={config.max_tokens}, temp={config.temperature})...")
tokens = generate(model, tokenizer, config, prompt_tokens)
logger.info(f"Generated {len(tokens)} tokens")
midi = tokenizer.tokens_to_midi(tokens)
output_path.parent.mkdir(parents=True, exist_ok=True)
midi.write(str(output_path))
logger.info(f"Saved MIDI to {output_path}")
return tokens