Spaces:
Running
Running
| <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> | |