| """ |
| 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() |
|
|
| |
| 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, |
| ) |
|
|
| |
| 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) |
|
|
| 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) |
| |
| 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.""" |
| |
| gen_config, prompt_tokens = parse_xml_config(xml_path) |
|
|
| |
| 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}") |
|
|
| |
| tokens = generate(model, tokenizer, gen_config, prompt_tokens, device=device) |
| |
| |
| 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}") |
|
|