AriaLM / run.py
krishnah27's picture
Upload folder using huggingface_hub
30e9297 verified
Raw
History Blame Contribute Delete
2.06 kB
#!/usr/bin/env python3
"""Standalone script: download dataset, train model, generate music."""
import subprocess
import sys
import os
import glob
BASE = os.path.dirname(os.path.abspath(__file__))
DATA_DIR = os.path.join(BASE, "data")
MIDI_DIR = os.path.join(DATA_DIR, "midi_files")
def download_dataset():
if os.path.exists(MIDI_DIR) and glob.glob(os.path.join(MIDI_DIR, "**", "*.mid"), recursive=True):
print(f"Dataset already exists at {MIDI_DIR}")
return
os.makedirs(DATA_DIR, exist_ok=True)
print("Downloading MIDI dataset via git clone...")
env = os.environ.copy()
env["GIT_LFS_SKIP_SMUDGE"] = "1"
result = subprocess.run(
["git", "clone", "--depth", "1",
"https://huggingface.co/datasets/drengskapur/midi-classical-music",
MIDI_DIR],
env=env,
capture_output=True,
text=True,
)
if result.returncode != 0:
print(f"Git clone failed: {result.stderr}")
# Fallback: use Python to download
print("Trying Python download fallback...")
from huggingface_hub import snapshot_download
snapshot_download(
repo_id="drengskapur/midi-classical-music",
repo_type="dataset",
local_dir=MIDI_DIR,
allow_patterns=["*.mid"],
)
midi_files = glob.glob(os.path.join(MIDI_DIR, "**", "*.mid"), recursive=True)
print(f"Downloaded {len(midi_files)} MIDI files")
def main():
# Step 1: Download
download_dataset()
# Remove stale cache
cache = os.path.join(DATA_DIR, "tokenized_cache.pkl")
if os.path.exists(cache):
os.remove(cache)
print("Removed stale cache")
# Step 2: Train + Generate
sys.path.insert(0, BASE)
os.chdir(BASE)
args = sys.argv[1:] if len(sys.argv) > 1 else ["train+generate", "--epochs", "20", "--batch-size", "4", "--seq-len", "512"]
sys.argv = ["s00_main"] + args
from src.s00_main import main as run_main
run_main()
if __name__ == "__main__":
main()