Allegretto-Mini-Demo / index.html
Dev4285's picture
Fix: enable WASM proxy for memory, add retry logic for chunk fetches
4b4bb3d verified
Raw
History Blame Contribute Delete
32.2 kB
<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="UTF-8" />
<meta name="viewport" content="width=device-width, initial-scale=1.0" />
<title>Allegretto Mini — AI Music Generation by OSAMA INC</title>
<link rel="preconnect" href="https://fonts.googleapis.com" />
<link rel="preconnect" href="https://fonts.gstatic.com" crossorigin />
<link href="https://fonts.googleapis.com/css2?family=Inter:wght@400;500;600;700&display=swap" rel="stylesheet" />
<style>
*, *::before, *::after { box-sizing: border-box; margin: 0; padding: 0; }
:root {
--bg: #0f0f0f; --surface: #1a1a2e; --surface-2: #16213e;
--accent: #e94560; --accent-hover: #ff6b81;
--text: #f1f1f1; --text-muted: #a0a0b0; --border: #2a2a40;
--success: #4ade80; --warning: #fbbf24; --radius: 12px;
}
body {
font-family: 'Inter', -apple-system, sans-serif;
background: var(--bg); color: var(--text);
min-height: 100vh; display: flex; flex-direction: column;
overflow-x: hidden;
}
.hero {
background: linear-gradient(135deg, #1a1a2e 0%, #16213e 40%, #0f3460 100%);
padding: 32px 24px 24px; text-align: center; border-bottom: 1px solid var(--border);
}
.hero-logo { font-size: 28px; font-weight: 700; letter-spacing: -0.5px; }
.hero-logo .accent { color: var(--accent); }
.hero-subtitle { margin-top: 8px; font-size: 14px; color: var(--text-muted); }
.hero-company {
display: inline-block; margin-top: 12px; padding: 4px 14px;
background: rgba(233,69,96,0.15); border: 1px solid rgba(233,69,96,0.3);
border-radius: 20px; font-size: 12px; font-weight: 600; color: var(--accent);
}
.main { flex: 1; max-width: 720px; width: 100%; margin: 0 auto; padding: 24px 16px; }
.prompt-section {
background: var(--surface); border: 1px solid var(--border);
border-radius: var(--radius); padding: 20px; margin-bottom: 20px;
}
.prompt-label { font-size: 14px; font-weight: 600; color: var(--text-muted); margin-bottom: 8px; }
.prompt-input {
width: 100%; min-height: 80px; padding: 12px 14px;
background: var(--bg); border: 1px solid var(--border);
border-radius: 8px; color: var(--text); font-family: 'Inter', sans-serif;
font-size: 15px; resize: vertical; outline: none; transition: border-color 0.2s;
}
.prompt-input:focus { border-color: var(--accent); }
.prompt-input::placeholder { color: var(--text-muted); opacity: 0.6; }
.settings-row { display: flex; gap: 16px; margin-bottom: 16px; flex-wrap: wrap; }
.setting-group { flex: 1; min-width: 120px; }
.setting-label { font-size: 12px; font-weight: 600; color: var(--text-muted); margin-bottom: 6px; }
.setting-input {
width: 100%; padding: 8px 12px; background: var(--bg);
border: 1px solid var(--border); border-radius: 8px;
color: var(--text); font-size: 14px; outline: none; transition: border-color 0.2s;
}
.setting-input:focus { border-color: var(--accent); }
.generate-btn {
display: flex; align-items: center; justify-content: center; gap: 8px;
width: 100%; padding: 14px; background: var(--accent); border: none;
border-radius: 8px; color: #fff; font-size: 15px; font-weight: 600;
cursor: pointer; transition: background 0.2s, transform 0.1s;
}
.generate-btn:hover { background: var(--accent-hover); }
.generate-btn:active { transform: scale(0.98); }
.generate-btn:disabled { background: #333; color: #666; cursor: not-allowed; }
.status-section {
background: var(--surface); border: 1px solid var(--border);
border-radius: var(--radius); padding: 16px 20px; margin-bottom: 20px;
}
.status-header { font-size: 14px; font-weight: 600; color: var(--text); margin-bottom: 10px; }
.progress-bar-track { width: 100%; height: 8px; background: var(--bg); border-radius: 4px; overflow: hidden; }
.progress-bar-fill { height: 100%; background: var(--accent); border-radius: 4px; transition: width 0.3s ease; width: 0%; }
.status-text { font-size: 13px; color: var(--text-muted); margin-top: 8px; }
.step-log {
margin-top: 10px; max-height: 200px; overflow-y: auto;
font-size: 12px; color: var(--text-muted); line-height: 1.6;
white-space: pre-wrap;
}
.step-log .step-ok { color: var(--success); }
.step-log .step-err { color: var(--accent); }
.output-section {
background: var(--surface); border: 1px solid var(--border);
border-radius: var(--radius); padding: 20px; margin-bottom: 20px;
}
.output-header { font-size: 14px; font-weight: 600; color: var(--text); margin-bottom: 12px; }
.audio-player { width: 100%; border-radius: 8px; outline: none; }
.download-link {
display: inline-flex; align-items: center; gap: 6px; margin-top: 12px;
padding: 8px 16px; background: var(--surface-2); border: 1px solid var(--border);
border-radius: 8px; color: var(--text); font-size: 13px; font-weight: 500;
cursor: pointer; text-decoration: none; transition: background 0.2s;
}
.download-link:hover { background: var(--border); }
.error-section {
background: rgba(233,69,96,0.1); border: 1px solid rgba(233,69,96,0.3);
border-radius: var(--radius); padding: 16px 20px; margin-bottom: 20px;
}
.error-text { font-size: 14px; color: var(--accent); word-break: break-word; }
.footer {
text-align: center; padding: 20px 16px; border-top: 1px solid var(--border);
font-size: 12px; color: var(--text-muted); line-height: 1.6;
}
.footer a { color: var(--accent); text-decoration: none; }
.footer a:hover { text-decoration: underline; }
@media (max-width: 480px) {
.hero { padding: 24px 16px 16px; }
.hero-logo { font-size: 22px; }
.main { padding: 16px 12px; }
.settings-row { gap: 10px; }
}
.spinner {
animation: spin 1s linear infinite; display: inline-block;
width: 18px; height: 18px; border: 2px solid rgba(255,255,255,0.3);
border-top-color: #fff; border-radius: 50%;
}
@keyframes spin { to { transform: rotate(360deg); } }
.step-log::-webkit-scrollbar { width: 4px; }
.step-log::-webkit-scrollbar-track { background: transparent; }
.step-log::-webkit-scrollbar-thumb { background: var(--border); border-radius: 2px; }
</style>
</head>
<body>
<header class="hero">
<div class="hero-logo">🎵 Allegretto <span class="accent">Mini</span></div>
<p class="hero-subtitle">Local AI music generation — text to stereo audio at 44.1 kHz</p>
<span class="hero-company">🇮🇳 OSAMA INC</span>
</header>
<main class="main">
<!-- Status - always visible -->
<section id="statusSection" class="status-section">
<div class="status-header" id="statusHeader">Initializing…</div>
<div class="progress-bar-track">
<div class="progress-bar-fill" id="progressBar"></div>
</div>
<div class="status-text" id="statusText">Loading ONNX Runtime Web</div>
<div class="step-log" id="stepLog"></div>
</section>
<!-- Error -->
<section id="errorSection" class="error-section" style="display:none">
<div class="error-text" id="errorText"></div>
</section>
<!-- Prompt -->
<section class="prompt-section" id="promptSection" style="display:none">
<label class="prompt-label" for="prompt">Describe the music you want to generate</label>
<textarea id="prompt" class="prompt-input" placeholder="e.g. upbeat electronic dance track with synth pads" rows="3">upbeat electronic dance track with synth pads</textarea>
<div class="settings-row">
<div class="setting-group">
<label class="setting-label" for="duration">Duration (seconds)</label>
<input type="number" id="duration" class="setting-input" value="10" min="1" max="30" step="1" />
</div>
<div class="setting-group">
<label class="setting-label" for="steps">Diffusion steps</label>
<input type="number" id="steps" class="setting-input" value="8" min="1" max="20" step="1" />
</div>
<div class="setting-group">
<label class="setting-label" for="seed">Seed</label>
<input type="number" id="seed" class="setting-input" value="0" min="0" max="999999" step="1" />
</div>
</div>
<button id="generateBtn" class="generate-btn">🎵 Generate Music</button>
</section>
<!-- Output -->
<section id="outputSection" class="output-section" style="display:none">
<div class="output-header">🎉 Generated Audio</div>
<audio id="audioPlayer" class="audio-player" controls></audio>
<a id="downloadLink" class="download-link" download="allegretto-mini-output.wav">💾 Download WAV</a>
</section>
</main>
<footer class="footer">
<strong>Allegretto Mini</strong> by <a href="https://github.com/aryanisproinroblox-source/Allegretto-Mini" target="_blank">OSAMA INC (India)</a><br />
Model on <a href="https://huggingface.co/Dev4285/Allegretto-Mini" target="_blank">HuggingFace</a> ·
Code on <a href="https://github.com/aryanisproinroblox-source/Allegretto-Mini" target="_blank">GitHub</a><br />
Runs entirely in your browser — no server, no cloud, no API keys needed.
</footer>
<!-- Load ONNX Runtime Web -->
<script src="https://cdn.jsdelivr.net/npm/onnxruntime-web@1.21.0/dist/ort.min.js"></script>
<script>
/* ═══════════════════════════════════════════
Allegretto Mini — Browser Music Generator
OSAMA INC (India)
═══════════════════════════════════════════ */
var MODEL_BASE = ".";
var SCHEDULE_CFG = { rate: 0, anchor_logsnr: -6.2, logsnr_end: 2.0 };
// ─── UI refs ───
var ui = {
prompt: document.getElementById("prompt"),
duration: document.getElementById("duration"),
steps: document.getElementById("steps"),
seed: document.getElementById("seed"),
generateBtn: document.getElementById("generateBtn"),
statusSection: document.getElementById("statusSection"),
statusHeader: document.getElementById("statusHeader"),
progressBar: document.getElementById("progressBar"),
statusText: document.getElementById("statusText"),
stepLog: document.getElementById("stepLog"),
errorSection: document.getElementById("errorSection"),
errorText: document.getElementById("errorText"),
promptSection: document.getElementById("promptSection"),
outputSection: document.getElementById("outputSection"),
audioPlayer: document.getElementById("audioPlayer"),
downloadLink: document.getElementById("downloadLink"),
};
// ─── State ───
var modelConfig = null;
var bpeTokenizer = null; // our custom BPE tokenizer object
var sessions = {}; // ONNX sessions
var modelLoaded = false;
// ─── Helpers ───
function log(msg, type) {
type = type || "ok";
var div = document.createElement("div");
div.className = "step-" + type;
div.textContent = msg;
ui.stepLog.appendChild(div);
ui.stepLog.scrollTop = ui.stepLog.scrollHeight;
}
function setStatus(header, text, pct) {
ui.statusHeader.textContent = header;
ui.statusText.textContent = text;
ui.progressBar.style.width = (pct || 0) + "%";
}
function showError(msg) {
ui.errorSection.style.display = "block";
ui.errorText.textContent = msg;
}
function randn(n, seed) {
var d = new Float32Array(n), s = seed;
for (var i = 0; i < n; i += 2) {
s = ((s * 1664525) + 1013904223) & 0xffffffff;
var u1 = (s >>> 0) / 0xffffffff || 1e-10;
s = ((s * 1664525) + 1013904223) & 0xffffffff;
var u2 = (s >>> 0) / 0xffffffff;
var r = Math.sqrt(-2.0 * Math.log(u1));
var th = 2.0 * Math.PI * u2;
d[i] = r * Math.cos(th);
if (i + 1 < n) d[i + 1] = r * Math.sin(th);
}
return d;
}
// ─── WAV Encoder ───
function encodeWAV(data, numFrames, sr) {
var nc = 2, bps = 2, ba = nc * bps, ds = numFrames * ba, bs = 44 + ds;
var buf = new ArrayBuffer(bs), v = new DataView(buf);
function ws(o, s) { for (var i = 0; i < s.length; i++) v.setUint8(o + i, s.charCodeAt(i)); }
ws(0, 'RIFF'); v.setUint32(4, bs - 8, true); ws(8, 'WAVE');
ws(12, 'fmt '); v.setUint32(16, 16, true); v.setUint16(20, 1, true);
v.setUint16(22, nc, true); v.setUint32(24, sr, true); v.setUint32(28, sr * ba, true);
v.setUint16(32, ba, true); v.setUint16(34, bps * 8, true);
ws(36, 'data'); v.setUint32(40, ds, true);
var off = 44;
for (var i = 0; i < numFrames; i++) {
for (var ch = 0; ch < nc; ch++) {
var idx = ch * numFrames + i;
var s = Math.max(-1, Math.min(1, data[idx]));
v.setInt16(off, s < 0 ? s * 0x8000 : s * 0x7FFF, true);
off += 2;
}
}
return buf;
}
// ─── BPE Schedule ───
function logsnrToTime(l) { return 1.0 / (1.0 + Math.exp(-l)); }
function computeSchedule(n, cfg) {
var t = [];
for (var i = 0; i <= n; i++) {
var f = i / n;
t.push(logsnrToTime(cfg.anchor_logsnr + f * (cfg.logsnr_end - cfg.anchor_logsnr)));
}
return t;
}
// ─── Fetch helpers ───
async function fetchJSON(path) {
var r = await fetch(MODEL_BASE + "/" + path);
if (!r.ok) throw new Error("Fetch failed: " + path + " (" + r.status + ")");
return r.json();
}
async function fetchBin(path, retries) {
retries = retries || 3;
for (var attempt = 0; attempt < retries; attempt++) {
try {
var r = await fetch(MODEL_BASE + "/" + path);
if (!r.ok) throw new Error("HTTP " + r.status);
return r.arrayBuffer();
} catch (e) {
if (attempt < retries - 1) {
log(" Retrying " + path + " (attempt " + (attempt + 2) + "/" + retries + ")…", "warn");
await new Promise(function(res) { setTimeout(res, 2000); });
} else {
throw new Error("Fetch failed after " + retries + " attempts: " + path + " — " + e.message);
}
}
}
}
/* ═══════════════════════════════════════════
BPE TOKENIZER — implemented from tokenizer.json
═══════════════════════════════════════════ */
function GemmaBPETokenizer(config) {
this.vocab = config.model.vocab; // { token_string: id }
this.merges = config.model.merges; // [[strA, strB], ...]
this.bosId = 2;
this.eosId = 1;
this.padId = 0;
this.unkId = 3;
// Build merge rank map for fast lookup
this.mergeRank = {};
for (var i = 0; i < this.merges.length; i++) {
var key = this.merges[i][0] + " " + this.merges[i][1];
this.mergeRank[key] = i;
}
// Build byte-fallback map: <0xNN> → token ID
this.byteFallback = {};
for (var k in this.vocab) {
if (k.length >= 4 && k[0] === '<' && k[1] === '0' && k[2] === 'x') {
var hex = k.slice(3, k.length - 1);
this.byteFallback[parseInt(hex, 16)] = this.vocab[k];
}
}
// Reverse vocab for debugging
this.idToToken = {};
for (var k in this.vocab) this.idToToken[this.vocab[k]] = k;
}
GemmaBPETokenizer.prototype.tokenize = function(text) {
// Gemma normalization: lowercase + NFKC-like
text = text.toLowerCase();
// Replace spaces with ▁ (SentencePiece convention)
var normalized = text.replace(/ /g, "▁");
// Pre-tokenize: split into words by ▁, keeping ▁ attached
// Each "word" starts with ▁ (except possibly the first if no leading space)
var words = [];
if (normalized.length > 0) {
// Split: ▁XXX patterns
var parts = normalized.split("▁");
for (var i = 0; i < parts.length; i++) {
if (parts[i].length > 0) {
if (i === 0 && normalized[0] !== "▁") {
words.push(parts[i]);
} else {
words.push("▁" + parts[i]);
}
}
}
}
// Apply BPE to each word
var tokens = [];
for (var w = 0; w < words.length; w++) {
var wordTokens = this.bpe(words[w]);
for (var t = 0; t < wordTokens.length; t++) {
tokens.push(wordTokens[t]);
}
}
// If no tokens, use UNK
if (tokens.length === 0) tokens.push(this.unkId);
// Add BOS + EOS
var result = [this.bosId];
for (var i = 0; i < tokens.length; i++) result.push(tokens[i]);
result.push(this.eosId);
return result;
};
GemmaBPETokenizer.prototype.bpe = function(word) {
// Split word into initial characters
var symbols = [];
for (var i = 0; i < word.length; i++) {
symbols.push(word[i]);
}
// Apply BPE merges iteratively
while (symbols.length > 1) {
// Find the pair with lowest merge rank
var bestRank = Infinity;
var bestIdx = -1;
for (var i = 0; i < symbols.length - 1; i++) {
var pairKey = symbols[i] + " " + symbols[i + 1];
var rank = this.mergeRank[pairKey];
if (rank !== undefined && rank < bestRank) {
bestRank = rank;
bestIdx = i;
}
}
if (bestIdx === -1) break; // No more merges possible
// Merge the best pair
var merged = symbols[bestIdx] + symbols[bestIdx + 1];
var newSymbols = [];
for (var i = 0; i < bestIdx; i++) newSymbols.push(symbols[i]);
newSymbols.push(merged);
for (var i = bestIdx + 2; i < symbols.length; i++) newSymbols.push(symbols[i]);
symbols = newSymbols;
}
// Convert symbols to IDs
var ids = [];
for (var i = 0; i < symbols.length; i++) {
var id = this.vocab[symbols[i]];
if (id !== undefined) {
ids.push(id);
} else {
// Byte fallback: encode each byte of the unknown symbol
var bytes = new TextEncoder().encode(symbols[i]);
for (var b = 0; b < bytes.length; b++) {
var byteId = this.byteFallback[bytes[b]];
if (byteId !== undefined) {
ids.push(byteId);
} else {
ids.push(this.unkId);
}
}
}
}
return ids;
};
/* ═══════════════════════════════════════════
LOAD ONNX SESSION
═══════════════════════════════════════════ */
async function loadSession(graphFile, manifestFile, progressBase, progressSpan) {
try {
setStatus("Loading model weights", "Fetching manifest " + manifestFile + "…", progressBase);
var manifest = await fetchJSON(manifestFile);
log(" Manifest fetched: " + (manifest.chunks ? manifest.chunks.length : 0) + " chunks");
setStatus("Loading model weights", "Fetching graph " + graphFile + "…", progressBase + 1);
var graphBuf = await fetchBin(graphFile);
log(" Graph fetched: " + (graphBuf.byteLength / 1048576).toFixed(1) + " MB");
var externalData = [];
var chunkBuffers = []; // Keep refs to free later
if (manifest.chunks && manifest.chunks.length > 0) {
var chunkDir = manifestFile.replace(/\/[^/]+$/, "");
for (var i = 0; i < manifest.chunks.length; i++) {
var name = manifest.chunks[i].name;
var size = manifest.chunks[i].size;
setStatus("Loading model weights", "Downloading " + name + " (" + (size / 1048576).toFixed(1) + " MB)…", progressBase + 2 + (i / manifest.chunks.length) * progressSpan);
try {
var buf = await fetchBin(chunkDir + "/" + name);
externalData.push({ path: name, data: new Uint8Array(buf) });
chunkBuffers.push(buf);
log(" Chunk " + name + " fetched: " + (buf.byteLength / 1048576).toFixed(1) + " MB");
} catch (chunkErr) {
log(" ✗ Chunk " + name + " failed: " + chunkErr.message, "err");
throw new Error("Chunk fetch failed: " + name + " — " + chunkErr.message);
}
}
}
var opts = { executionProviders: ["wasm"] };
if (externalData.length > 0) opts.externalData = externalData;
setStatus("Loading model weights", "Creating ONNX session for " + graphFile + "…", progressBase + progressSpan);
log(" Creating session (externalData: " + externalData.length + " entries)…");
var session = await ort.InferenceSession.create(new Uint8Array(graphBuf), opts);
log(" Session created! Input names: " + session.inputNames.join(", "));
// Free download buffers to reclaim memory before loading next model
externalData = null;
chunkBuffers = null;
graphBuf = null;
// Brief pause to allow GC
await new Promise(function(r) { setTimeout(r, 500); });
return session;
} catch (err) {
log(" ✗ loadSession failed for " + graphFile + ": " + err.message, "err");
throw err;
}
}
/* ═══════════════════════════════════════════
INITIALIZATION
═══════════════════════════════════════════ */
async function initApp() {
try {
if (typeof ort === "undefined") {
throw new Error("ONNX Runtime Web failed to load. Try refreshing the page.");
}
setStatus("Initializing", "ONNX Runtime Web loaded ✓", 2);
log("✓ ONNX Runtime Web loaded");
// Configure WASM with larger memory for all model weights (~640MB)
ort.env.wasm.numThreads = 1;
ort.env.wasm.simd = true;
ort.env.wasm.proxy = true; // Use proxy worker for better memory handling
log("✓ WASM configured (1 thread, SIMD on)");
// Load config
setStatus("Loading config", "Fetching model config…", 4);
modelConfig = await fetchJSON("config.json");
log("✓ Config loaded — " + modelConfig.sample_rate + " Hz, " + modelConfig.default_seconds + "s default");
// Load tokenizer
setStatus("Loading tokenizer", "Fetching tokenizer.json (~34 MB)…", 6);
var tokenizerData = await fetchJSON("tokenizer/tokenizer.json");
bpeTokenizer = new GemmaBPETokenizer(tokenizerData);
log("✓ BPE tokenizer loaded — " + Object.keys(bpeTokenizer.vocab).length + " vocab entries");
// Quick test tokenization
var testIds = bpeTokenizer.tokenize("hello world");
log("✓ Test tokenize 'hello world': " + JSON.stringify(testIds.slice(0, 10)) + "… (" + testIds.length + " tokens)");
// Load ONNX sessions — smallest first, free memory after each
setStatus("Loading model weights", "Loading number conditioner (~0.8 MB)…", 10);
sessions.numberConditioner = await loadSession("onnx/number_conditioner.onnx", "onnx/number_conditioner_chunks.json", 10, 5);
log("✓ Number conditioner loaded");
setStatus("Loading model weights", "Loading decoder (~45 MB)…", 15);
sessions.decoder = await loadSession("onnx/decoder_q4.onnx", "onnx/decoder_q4_chunks.json", 15, 10);
log("✓ Decoder loaded");
setStatus("Loading model weights", "Loading text encoder (~213 MB)…", 25);
sessions.textEncoder = await loadSession("onnx/text_encoder_q4.onnx", "onnx/text_encoder_q4_chunks.json", 25, 25);
log("✓ Text encoder loaded");
setStatus("Loading model weights", "Loading DiT transformer (~380 MB)…", 50);
sessions.dit = await loadSession("onnx/dit_q4.onnx", "onnx/dit_q4_chunks.json", 50, 45);
log("✓ DiT loaded — all model weights ready!");
modelLoaded = true;
setStatus("Ready! 🎵", "Model loaded. Enter a prompt and click Generate.", 100);
ui.promptSection.style.display = "block";
} catch (err) {
console.error("Init error:", err);
showError("Initialization failed: " + err.message);
log("✗ " + err.message, "err");
setStatus("Error", err.message, 0);
}
}
/* ═══════════════════════════════════════════
GENERATION PIPELINE
═══════════════════════════════════════════ */
async function generate() {
var prompt = ui.prompt.value.trim();
if (!prompt) { showError("Enter a music prompt first."); return; }
var seconds = parseInt(ui.duration.value) || 10;
var numSteps = parseInt(ui.steps.value) || 8;
var seed = parseInt(ui.seed.value) || 0;
ui.outputSection.style.display = "none";
ui.errorSection.style.display = "none";
ui.generateBtn.disabled = true;
ui.generateBtn.innerHTML = '<span class="spinner"></span> Generating…';
ui.stepLog.innerHTML = "";
try {
// ─── Tokenize ───
setStatus("Encoding prompt", "Tokenizing…", 60);
var ids = bpeTokenizer.tokenize(prompt);
// Pad/truncate to model's expected length (256)
var maxLen = modelConfig.text_max_length;
if (ids.length > maxLen) ids = ids.slice(0, maxLen);
// Right-pad with pad_id (0) to reach maxLen
while (ids.length < maxLen) ids.push(bpeTokenizer.padId);
var inputIdsArr = new BigInt64Array(ids.length);
var attnMaskArr = new BigInt64Array(ids.length);
// attention_mask: 1 for real tokens (up to original length), 0 for padding
var origLen = bpeTokenizer.tokenize(prompt).length;
for (var i = 0; i < ids.length; i++) {
inputIdsArr[i] = BigInt(ids[i]);
attnMaskArr[i] = BigInt(i < origLen ? 1 : 0);
}
var inputIds = new ort.Tensor("int64", inputIdsArr, [1, ids.length]);
var attentionMask = new ort.Tensor("int64", attnMaskArr, [1, ids.length]);
log("✓ Tokenized: " + ids.length + " tokens from '" + prompt + "'");
// ─── Text Encoder ───
setStatus("Running text encoder", "Encoding text…", 65);
var encOut = await sessions.textEncoder.run({ input_ids: inputIds, attention_mask: attentionMask });
var textEmb = encOut.last_hidden_state;
log("✓ Text embeddings: [" + textEmb.dims.join(",") + "]");
// ─── Duration Embedding ───
setStatus("Computing duration", seconds + "s…", 70);
var secTensor = new ort.Tensor("float32", new Float32Array([seconds]), [1]);
var ncOut = await sessions.numberConditioner.run({ seconds: secTensor });
var durEmb = ncOut.embedding;
// Reshape from [1,1,768] to [1,768] — global_embed expects rank 2
if (durEmb.dims.length === 3 && durEmb.dims[1] === 1) {
var reshapedData = new Float32Array(durEmb.data.length);
for (var i = 0; i < durEmb.data.length; i++) reshapedData[i] = durEmb.data[i];
durEmb = new ort.Tensor("float32", reshapedData, [1, durEmb.dims[2]]);
}
log("✓ Duration embedding: [" + durEmb.dims.join(",") + "]");
// ─── Build Conditioning ───
var seqLen = textEmb.dims[1];
var condDim = modelConfig.cond_dim;
var crossAttnData = new Float32Array((seqLen + 1) * condDim);
crossAttnData.set(textEmb.data, 0);
crossAttnData.set(durEmb.data, seqLen * condDim);
var crossAttn = new ort.Tensor("float32", crossAttnData, [1, seqLen + 1, condDim]);
var globalCond = durEmb;
// ─── Latent shape ───
var T_lat = Math.ceil((seconds + 6) * modelConfig.sample_rate / modelConfig.audio_align) * 2;
var latentShape = [1, modelConfig.io_channels, T_lat];
log("✓ Latent: [1," + modelConfig.io_channels + "," + T_lat + "]");
// ─── Sampling Loop ───
setStatus("Generating music", numSteps + " diffusion steps…", 75);
var schedule = computeSchedule(numSteps + 1, SCHEDULE_CFG);
var totalElts = latentShape[0] * latentShape[1] * latentShape[2];
var noise = randn(totalElts, seed);
var x = new ort.Tensor("float32", noise, latentShape);
for (var step = 0; step < numSteps; step++) {
var t_curr = schedule[step];
var t_next = schedule[step + 1];
var tTensor = new ort.Tensor("float32", new Float32Array([t_curr]), [1]);
var padMask = new ort.Tensor("bool", new Uint8Array(T_lat).fill(1), [1, T_lat]);
var crossPadMask = new ort.Tensor("bool", new Uint8Array(seqLen + 1).fill(0), [1, seqLen + 1]);
var localAdd = new ort.Tensor("float32", new Float32Array((seqLen + 1) * T_lat), [1, seqLen + 1, T_lat]);
var ditOut = await sessions.dit.run({
x: x, t: tTensor,
cross_attn_cond: crossAttn,
global_embed: globalCond,
local_add_cond: localAdd,
padding_mask: padMask,
});
// Pingpong: denoised = x - t * dit_out; x_next = (1-t_next)*denoised
var ditResult = ditOut.out;
var xNext = new Float32Array(totalElts);
for (var i = 0; i < totalElts; i++) {
var denoised = x.data[i] - t_curr * ditResult.data[i];
xNext[i] = (1.0 - t_next) * denoised;
}
x = new ort.Tensor("float32", xNext, latentShape);
var pct = 75 + ((step + 1) / numSteps) * 15;
setStatus("Generating music", "Step " + (step + 1) + "/" + numSteps, pct);
log(" Step " + (step + 1) + "/" + numSteps + " (t=" + t_curr.toFixed(3) + "→" + t_next.toFixed(3) + ")");
}
log("✓ Diffusion done");
// ─── Decode ───
setStatus("Decoding audio", "Running decoder…", 92);
var decPadMask = new ort.Tensor("bool", new Uint8Array(T_lat).fill(1), [1, T_lat]);
var decOut = await sessions.decoder.run({ latents: x });
var audio = decOut.audio;
log("✓ Audio decoded: [" + audio.dims.join(",") + "]");
// ─── WAV ───
setStatus("Encoding WAV…", "Creating audio file…", 96);
var numFrames = audio.dims[2];
var wavBuf = encodeWAV(audio.data, numFrames, modelConfig.sample_rate);
var wavBlob = new Blob([wavBuf], { type: "audio/wav" });
ui.audioPlayer.src = URL.createObjectURL(wavBlob);
ui.downloadLink.href = URL.createObjectURL(wavBlob);
ui.downloadLink.download = "allegretto-mini-" + seconds + "s.wav";
ui.outputSection.style.display = "block";
setStatus("Done! 🎉", "Music generated successfully.", 100);
log("✓ WAV ready — play or download!");
} catch (err) {
console.error("Gen error:", err);
showError("Generation failed: " + err.message);
log("✗ " + err.message, "err");
setStatus("Error", err.message, 0);
} finally {
ui.generateBtn.disabled = false;
ui.generateBtn.innerHTML = "🎵 Generate Music";
}
}
// ─── Boot ───
ui.generateBtn.addEventListener("click", generate);
initApp();
</script>
</body>
</html>