""" XML Configuration-driven Inference for AriaLM. Allows specifying prompt notes, generation length, and sampling settings in an XML file. """ import logging import xml.etree.ElementTree as ET from pathlib import Path import pretty_midi import torch from src.s01_config import GenConfig, ModelConfig, PathConfig from src.s02_tokenizer import MusicTokenizer from src.s04_model import MusicTransformer from src.s06_generator import generate logger = logging.getLogger(__name__) def parse_xml_config(xml_path: Path) -> tuple[GenConfig, list[int]]: """ Parse generation settings and starting prompt notes from an XML file. """ tree = ET.parse(xml_path) root = tree.getroot() # 1. Parse Settings settings = root.find("settings") temp = float(settings.find("temperature").text or 0.85) top_k = int(settings.find("top_k").text or 40) top_p = float(settings.find("top_p").text or 0.92) max_tokens = int(settings.find("max_tokens").text or 512) rep_penalty = float(settings.find("repetition_penalty").text or 1.15) seed = int(settings.find("seed").text or 42) gen_config = GenConfig( temperature=temp, top_k=top_k, top_p=top_p, max_tokens=max_tokens, repetition_penalty=rep_penalty, seed=seed, ) # 2. Parse Prompt Notes prompt_node = root.find("prompt") prompt_tokens = [] if prompt_node is not None: tokenizer = MusicTokenizer() pm = pretty_midi.PrettyMIDI() instrument = pretty_midi.Instrument(program=0) # Default piano current_time = 0.0 for note_el in prompt_node.findall("note"): pitch = int(note_el.attrib["pitch"]) velocity = int(note_el.attrib["velocity"]) duration = float(note_el.attrib["duration_ms"]) / 1000.0 delay = float(note_el.attrib.get("delay_ms", 0)) / 1000.0 current_time += delay note = pretty_midi.Note( velocity=velocity, pitch=pitch, start=current_time, end=current_time + duration, ) instrument.notes.append(note) current_time += duration pm.instruments.append(instrument) # Convert prompt MIDI to token IDs prompt_tokens = tokenizer.midi_to_tokens(pm) return gen_config, prompt_tokens def generate_from_xml(xml_path: Path, model_path: Path, output_path: Path): """Run generation using parameters specified in an XML file.""" # Load config and prompt gen_config, prompt_tokens = parse_xml_config(xml_path) # Load Model tokenizer = MusicTokenizer() model_config = ModelConfig(vocab_size=tokenizer.vocab_size) model = MusicTransformer.from_config(model_config) ckpt = torch.load(model_path, map_location="cpu", weights_only=False) model.load_state_dict(ckpt["model_state_dict"]) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model.to(device) print(f"Generating from XML Config: {xml_path}") print(f"Sampling details: Temp={gen_config.temperature}, Seed={gen_config.seed}, Length={gen_config.max_tokens}") # Generate tokens = generate(model, tokenizer, gen_config, prompt_tokens, device=device) # Save output midi = tokenizer.tokens_to_midi(tokens) output_path.parent.mkdir(parents=True, exist_ok=True) midi.write(str(output_path)) print(f"Saved generated MIDI to: {output_path}")