From d32dc35b6dfd87abaf07d2527e341148c0455272 Mon Sep 17 00:00:00 2001 From: Xiaohan-Tian <157918347+Xiaohan-Tian@users.noreply.github.com> Date: Mon, 11 May 2026 22:27:22 -0700 Subject: [PATCH] feat: added local separator implementation --- package-lock.json | 139 ++++- package.json | 1 + src/components/KGOnePanel.css | 95 ++- src/components/KGOnePanel.test.tsx | 180 ++++++ src/components/KGOnePanel.tsx | 396 ++++++++++++- src/components/Toolbar.tsx | 6 +- src/core/io/LocalSeparatorModelCache.test.ts | 151 +++++ src/util/localSeparatorConfig.ts | 31 + src/util/localSeparatorCpuDsp.ts | 139 +++++ src/util/localSeparatorGpuDsp.ts | 220 +++++++ src/util/localSeparatorModelCache.ts | 126 ++++ src/util/localSeparatorRunner.ts | 589 +++++++++++++++++++ src/util/localSeparatorRuntime.ts | 95 +++ src/util/localSeparatorShared.ts | 275 +++++++++ src/util/localSeparatorTiming.ts | 52 ++ src/util/localSeparatorTypes.ts | 49 ++ 16 files changed, 2506 insertions(+), 38 deletions(-) create mode 100644 src/components/KGOnePanel.test.tsx create mode 100644 src/core/io/LocalSeparatorModelCache.test.ts create mode 100644 src/util/localSeparatorConfig.ts create mode 100644 src/util/localSeparatorCpuDsp.ts create mode 100644 src/util/localSeparatorGpuDsp.ts create mode 100644 src/util/localSeparatorModelCache.ts create mode 100644 src/util/localSeparatorRunner.ts create mode 100644 src/util/localSeparatorRuntime.ts create mode 100644 src/util/localSeparatorShared.ts create mode 100644 src/util/localSeparatorTiming.ts create mode 100644 src/util/localSeparatorTypes.ts diff --git a/package-lock.json b/package-lock.json index 3ce903b..bc951f5 100644 --- a/package-lock.json +++ b/package-lock.json @@ -1,18 +1,19 @@ { "name": "K.G.Studio", - "version": "0.15.0-build.20260510", + "version": "0.16.0-build.20260510", "lockfileVersion": 3, "requires": true, "packages": { "": { "name": "K.G.Studio", - "version": "0.15.0-build.20260510", + "version": "0.16.0-build.20260510", "dependencies": { "@breezystack/lamejs": "^1.2.7", "class-transformer": "^0.5.1", "fft.js": "^4.0.4", "idb": "^8.0.3", "jszip": "^3.10.1", + "onnxruntime-web": "^1.26.0", "openai": "^6.33.0", "react": "^19.1.0", "react-dom": "^19.1.0", @@ -1628,6 +1629,70 @@ "dev": true, "license": "MIT" }, + "node_modules/@protobufjs/aspromise": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/@protobufjs/aspromise/-/aspromise-1.1.2.tgz", + "integrity": "sha512-j+gKExEuLmKwvz3OgROXtrJ2UG2x8Ch2YZUxahh+s1F2HZ+wAceUNLkvy6zKCPVRkU++ZWQrdxsUeQXmcg4uoQ==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/base64": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/@protobufjs/base64/-/base64-1.1.2.tgz", + "integrity": "sha512-AZkcAA5vnN/v4PDqKyMR5lx7hZttPDgClv83E//FMNhR2TMcLUhfRUBHCmSl0oi9zMgDDqRUJkSxO3wm85+XLg==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/codegen": { + "version": "2.0.5", + "resolved": "https://registry.npmjs.org/@protobufjs/codegen/-/codegen-2.0.5.tgz", + "integrity": "sha512-zgXFLzW3Ap33e6d0Wlj4MGIm6Ce8O89n/apUaGNB/jx+hw+ruWEp7EwGUshdLKVRCxZW12fp9r40E1mQrf/34g==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/eventemitter": { + "version": "1.1.0", + "resolved": "https://registry.npmjs.org/@protobufjs/eventemitter/-/eventemitter-1.1.0.tgz", + "integrity": "sha512-j9ednRT81vYJ9OfVuXG6ERSTdEL1xVsNgqpkxMsbIabzSo3goCjDIveeGv5d03om39ML71RdmrGNjG5SReBP/Q==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/fetch": { + "version": "1.1.0", + "resolved": "https://registry.npmjs.org/@protobufjs/fetch/-/fetch-1.1.0.tgz", + "integrity": "sha512-lljVXpqXebpsijW71PZaCYeIcE5on1w5DlQy5WH6GLbFryLUrBD4932W/E2BSpfRJWseIL4v/KPgBFxDOIdKpQ==", + "license": "BSD-3-Clause", + "dependencies": { + "@protobufjs/aspromise": "^1.1.1", + "@protobufjs/inquire": "^1.1.0" + } + }, + "node_modules/@protobufjs/float": { + "version": "1.0.2", + "resolved": "https://registry.npmjs.org/@protobufjs/float/-/float-1.0.2.tgz", + "integrity": "sha512-Ddb+kVXlXst9d+R9PfTIxh1EdNkgoRe5tOX6t01f1lYWOvJnSPDBlG241QLzcyPdoNTsblLUdujGSE4RzrTZGQ==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/inquire": { + "version": "1.1.1", + "resolved": "https://registry.npmjs.org/@protobufjs/inquire/-/inquire-1.1.1.tgz", + "integrity": "sha512-mnzgDV26ueAvk7rsbt9L7bE0SuAoqyuys/sMMrmVcN5x9VsxpcG3rqAUSgDyLp0UZlmNfIbQ4fHfCtreVBk8Ew==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/path": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/@protobufjs/path/-/path-1.1.2.tgz", + "integrity": "sha512-6JOcJ5Tm08dOHAbdR3GrvP+yUUfkjG5ePsHYczMFLq3ZmMkAD98cDgcT2iA1lJ9NVwFd4tH/iSSoe44YWkltEA==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/pool": { + "version": "1.1.0", + "resolved": "https://registry.npmjs.org/@protobufjs/pool/-/pool-1.1.0.tgz", + "integrity": "sha512-0kELaGSIDBKvcgS4zkjz1PeddatrjYcmMWOlAuAPwAeccUrPHdUqo/J6LiymHHEiJT5NrF1UVwxY14f+fy4WQw==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/utf8": { + "version": "1.1.1", + "resolved": "https://registry.npmjs.org/@protobufjs/utf8/-/utf8-1.1.1.tgz", + "integrity": "sha512-oOAWABowe8EAbMyWKM0tYDKi8Yaox52D+HWZhAIJqQXbqe0xI/GV7FhLWqlEKreMkfDjshR5FKgi3mnle0h6Eg==", + "license": "BSD-3-Clause" + }, "node_modules/@rolldown/pluginutils": { "version": "1.0.0-beta.19", "resolved": "https://registry.npmjs.org/@rolldown/pluginutils/-/pluginutils-1.0.0-beta.19.tgz", @@ -2630,7 +2695,6 @@ "version": "22.18.0", "resolved": "https://registry.npmjs.org/@types/node/-/node-22.18.0.tgz", "integrity": "sha512-m5ObIqwsUp6BZzyiy4RdZpzWGub9bqLJMvZDD0QMXhxjqMHMENlj+SqF5QxoUwaQNFe+8kz8XM8ZQhqkQPTgMQ==", - "dev": true, "license": "MIT", "dependencies": { "undici-types": "~6.21.0" @@ -4802,6 +4866,12 @@ "node": ">=16" } }, + "node_modules/flatbuffers": { + "version": "25.9.23", + "resolved": "https://registry.npmjs.org/flatbuffers/-/flatbuffers-25.9.23.tgz", + "integrity": "sha512-MI1qs7Lo4Syw0EOzUl0xjs2lsoeqFku44KpngfIduHBYvzm8h2+7K8YMQh1JtVVVrUvhLpNwqVi4DERegUJhPQ==", + "license": "Apache-2.0" + }, "node_modules/flatted": { "version": "3.3.3", "resolved": "https://registry.npmjs.org/flatted/-/flatted-3.3.3.tgz", @@ -5087,6 +5157,12 @@ "dev": true, "license": "MIT" }, + "node_modules/guid-typescript": { + "version": "1.0.9", + "resolved": "https://registry.npmjs.org/guid-typescript/-/guid-typescript-1.0.9.tgz", + "integrity": "sha512-Y8T4vYhEfwJOTbouREvG+3XDsjr8E3kIr7uf+JZ0BYloFsttiHU0WfvANVsR7TxNUJa/WpCnw/Ino/p+DeBhBQ==", + "license": "ISC" + }, "node_modules/handlebars": { "version": "4.7.8", "resolved": "https://registry.npmjs.org/handlebars/-/handlebars-4.7.8.tgz", @@ -6121,6 +6197,12 @@ "dev": true, "license": "MIT" }, + "node_modules/long": { + "version": "5.3.2", + "resolved": "https://registry.npmjs.org/long/-/long-5.3.2.tgz", + "integrity": "sha512-mNAgZ1GmyNhD7AuqnTG3/VQ26o760+ZYBPKjPvugO8+nLbYfX6TVpJPseBvopbdY+qpZ/lKUnmEc1LeZYS3QAA==", + "license": "Apache-2.0" + }, "node_modules/longest-streak": { "version": "3.1.0", "resolved": "https://registry.npmjs.org/longest-streak/-/longest-streak-3.1.0.tgz", @@ -10920,6 +11002,26 @@ "url": "https://github.com/sponsors/sindresorhus" } }, + "node_modules/onnxruntime-common": { + "version": "1.26.0", + "resolved": "https://registry.npmjs.org/onnxruntime-common/-/onnxruntime-common-1.26.0.tgz", + "integrity": "sha512-qVyMR4lcWgbkc4getFV+GQijsTnbg/siteoqcDwa3sI/LxbrMSNw4ePyvCq/ymdQaRomCA7YuWmhzsswxvymdw==", + "license": "MIT" + }, + "node_modules/onnxruntime-web": { + "version": "1.26.0", + "resolved": "https://registry.npmjs.org/onnxruntime-web/-/onnxruntime-web-1.26.0.tgz", + "integrity": "sha512-LbRr/8zZt2xilI2smrVQGGKINo0U46i8qJp+UXyMBGfqN7KjnH1BiwCwLwyNIVV4i9CKFv7Sf4PwLKWnT8/bEA==", + "license": "MIT", + "dependencies": { + "flatbuffers": "^25.1.24", + "guid-typescript": "^1.0.9", + "long": "^5.2.3", + "onnxruntime-common": "1.26.0", + "platform": "^1.3.6", + "protobufjs": "^7.2.4" + } + }, "node_modules/openai": { "version": "6.33.0", "resolved": "https://registry.npmjs.org/openai/-/openai-6.33.0.tgz", @@ -11341,6 +11443,12 @@ "node": ">=4" } }, + "node_modules/platform": { + "version": "1.3.6", + "resolved": "https://registry.npmjs.org/platform/-/platform-1.3.6.tgz", + "integrity": "sha512-fnWVljUchTro6RiCFvCXBbNhJc2NijN7oIQxbwsyL0buWJPG85v81ehlHI9fXrJsMNgTofEoWIQeClKpgxFLrg==", + "license": "MIT" + }, "node_modules/postcss": { "version": "8.5.6", "resolved": "https://registry.npmjs.org/postcss/-/postcss-8.5.6.tgz", @@ -11442,6 +11550,30 @@ "dev": true, "license": "ISC" }, + "node_modules/protobufjs": { + "version": "7.5.7", + "resolved": "https://registry.npmjs.org/protobufjs/-/protobufjs-7.5.7.tgz", + "integrity": "sha512-NGnrxS/nLKUo5nkbVQxlC71sB4hdfImdYIbFeSCidxtwATx0AHRPcANSLd0q5Bb2BkoSWo2iisQhGg5/r+ihbA==", + "hasInstallScript": true, + "license": "BSD-3-Clause", + "dependencies": { + "@protobufjs/aspromise": "^1.1.2", + "@protobufjs/base64": "^1.1.2", + "@protobufjs/codegen": "^2.0.5", + "@protobufjs/eventemitter": "^1.1.0", + "@protobufjs/fetch": "^1.1.0", + "@protobufjs/float": "^1.0.2", + "@protobufjs/inquire": "^1.1.1", + "@protobufjs/path": "^1.1.2", + "@protobufjs/pool": "^1.1.0", + "@protobufjs/utf8": "^1.1.1", + "@types/node": ">=13.7.0", + "long": "^5.0.0" + }, + "engines": { + "node": ">=12.0.0" + } + }, "node_modules/punycode": { "version": "2.3.1", "resolved": "https://registry.npmjs.org/punycode/-/punycode-2.3.1.tgz", @@ -13531,7 +13663,6 @@ "version": "6.21.0", "resolved": "https://registry.npmjs.org/undici-types/-/undici-types-6.21.0.tgz", "integrity": "sha512-iwDZqg0QAGrg9Rav5H4n0M64c3mkR59cJ6wQp+7C4nI0gsmExaedaYLNO44eT4AtBBwjbTiGPMlt2Md0T9H9JQ==", - "dev": true, "license": "MIT" }, "node_modules/unicorn-magic": { diff --git a/package.json b/package.json index 8ab1ffb..68de2aa 100644 --- a/package.json +++ b/package.json @@ -19,6 +19,7 @@ "fft.js": "^4.0.4", "idb": "^8.0.3", "jszip": "^3.10.1", + "onnxruntime-web": "^1.26.0", "openai": "^6.33.0", "react": "^19.1.0", "react-dom": "^19.1.0", diff --git a/src/components/KGOnePanel.css b/src/components/KGOnePanel.css index d37a191..d7d79c9 100644 --- a/src/components/KGOnePanel.css +++ b/src/components/KGOnePanel.css @@ -62,6 +62,15 @@ border-bottom-color: #5a9fd4; } +.kgone-tab.is-disabled { + color: #666; + cursor: not-allowed; +} + +.kgone-tab.is-disabled:hover { + color: #666; +} + /* Scrollable body */ .kgone-panel-body { flex: 1; @@ -203,6 +212,64 @@ border-radius: 6px; } +.kgone-local-mode-card { + background-color: #252525; + border: 1px solid #3a3a3a; + border-radius: 6px; + padding: 10px; + display: flex; + flex-direction: column; + gap: 8px; +} + +.kgone-local-mode-title { + color: #e0e0e0; + font-size: 12px; + font-weight: 600; +} + +.kgone-local-mode-text { + color: #999; + font-size: 11px; + line-height: 1.45; +} + +.kgone-runtime-row { + display: flex; + flex-wrap: wrap; + gap: 6px; +} + +.kgone-provider-chip { + background-color: #313131; + border: 1px solid #444; + border-radius: 999px; + color: #bbb; + font-size: 10px; + padding: 4px 8px; +} + +.kgone-progress-block { + display: flex; + flex-direction: column; + gap: 6px; +} + +.kgone-progress-track { + width: 100%; + height: 8px; + background-color: #1d1d1d; + border: 1px solid #3a3a3a; + border-radius: 999px; + overflow: hidden; +} + +.kgone-progress-fill { + height: 100%; + background: linear-gradient(90deg, #5a9fd4 0%, #7cc2f1 100%); + transition: width 0.15s linear; +} + .kgone-region-info { background-color: #252525; border: 1px solid #3a3a3a; @@ -252,6 +319,32 @@ cursor: not-allowed; } +.kgone-btn-secondary { + background-color: #353535; + border: 1px solid #4a4a4a; + border-radius: 4px; + color: #ddd; + font-size: 11px; + font-weight: 600; + padding: 8px 10px; + cursor: pointer; + flex: 1; +} + +.kgone-btn-secondary:hover:not(:disabled) { + background-color: #404040; +} + +.kgone-btn-secondary:disabled { + color: #666; + cursor: not-allowed; +} + +.kgone-btn-danger { + border-color: #5a2a2a; + color: #f0b0b0; +} + /* ── Audio player ──────────────────────────────────────────────────────────── */ .kgone-audio-player { display: flex; @@ -404,4 +497,4 @@ font-size: 11px; cursor: pointer; user-select: none; -} \ No newline at end of file +} diff --git a/src/components/KGOnePanel.test.tsx b/src/components/KGOnePanel.test.tsx new file mode 100644 index 0000000..f63f5bf --- /dev/null +++ b/src/components/KGOnePanel.test.tsx @@ -0,0 +1,180 @@ +import React from 'react'; +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { fireEvent, render, screen, waitFor } from '@testing-library/react'; +import KGOnePanel from './KGOnePanel'; +import { KGAudioRegion } from '../core/region/KGAudioRegion'; +import { KGAudioTrack } from '../core/track/KGAudioTrack'; + +let kgoneEnabled = false; +let selectedRegionIds: string[] = []; +let localModelCached = false; +let localSeparationResult: Array<{ name: string; blob: Blob }> = []; + +const mockRefreshProjectState = vi.fn(); +const mockExecuteCommand = vi.fn(); + +vi.mock('../stores/projectStore', () => ({ + useProjectStore: () => ({ + selectedRegionIds, + projectName: 'Test Project', + bpm: 120, + keySignature: 'C major', + timeSignature: { numerator: 4, denominator: 4 }, + maxBars: 32, + refreshProjectState: mockRefreshProjectState, + }), +})); + +vi.mock('../core/config/ConfigManager', () => ({ + ConfigManager: { + instance: () => ({ + get: (key: string) => { + if (key === 'general.kgone.enabled') return kgoneEnabled; + if (key === 'general.kgone.base_url') return 'http://127.0.0.1:8000'; + return undefined; + }, + }), + }, +})); + +const audioRegion = new KGAudioRegion( + 'audio-region-1', + 'track-1', + 0, + 'Verse Stem', + 0, + 4, + 'audio-file-1', + 'verse.wav', + 2, + 0, +); +const audioTrack = new KGAudioTrack('Audio Track', 1); +audioTrack.setTrackIndex(0); +audioTrack.setRegions([audioRegion]); + +vi.mock('../core/KGCore', () => ({ + KGCore: { + instance: () => ({ + getCurrentProject: () => ({ + getTracks: () => [audioTrack], + getName: () => 'Test Project', + }), + executeCommand: mockExecuteCommand, + }), + }, +})); + +vi.mock('../core/io/KGAudioFileStorage', () => ({ + KGAudioFileStorage: { + loadAudioFile: vi.fn(async () => new ArrayBuffer(8)), + storeAudioFile: vi.fn(async () => undefined), + }, +})); + +vi.mock('../util/audioUtil', () => ({ + sliceAudioToWav: vi.fn(async (_buffer: ArrayBuffer) => _buffer), +})); + +vi.mock('../util/localSeparatorModelCache', () => ({ + LocalSeparatorModelCache: { + exists: vi.fn(async () => localModelCached), + download: vi.fn(async () => { + localModelCached = true; + }), + delete: vi.fn(async () => { + localModelCached = false; + }), + getArrayBuffer: vi.fn(async () => new ArrayBuffer(16)), + }, +})); + +vi.mock('../util/localSeparatorRuntime', () => ({ + detectLocalRuntimeSupport: () => ({ webgpuExposed: false }), + LocalOrtRuntimeManager: class { + constructor(private readonly options?: { onProviderChange?: (provider: string) => void }) {} + + reset() {} + + async ensureRuntime() { + this.options?.onProviderChange?.('cpu/wasm'); + return { provider: 'wasm', session: {} }; + } + }, +})); + +vi.mock('../util/localSeparatorRunner', () => ({ + runLocalSeparator: vi.fn(async ({ onProgress, onProviderChange }) => { + onProviderChange?.('cpu/wasm'); + onProgress({ stage: 'main', passLabel: 'Main pass', percent: 100, processedChunks: 1, totalChunks: 1 }); + return { + stems: localSeparationResult, + providerLabel: 'CPU/wasm', + debugSummary: {}, + }; + }), +})); + +describe('KGOnePanel local separator mode', () => { + beforeEach(() => { + kgoneEnabled = false; + selectedRegionIds = []; + localModelCached = false; + localSeparationResult = [ + { name: 'Instrumental', blob: new Blob(['instrumental'], { type: 'audio/wav' }) }, + { name: 'Vocals', blob: new Blob(['vocals'], { type: 'audio/wav' }) }, + ]; + mockRefreshProjectState.mockReset(); + mockExecuteCommand.mockReset(); + }); + + it('defaults to separator and shows other tabs as disabled in local mode', async () => { + render(); + + expect(await screen.findByText('Local Separator Mode')).toBeInTheDocument(); + expect(screen.getByRole('button', { name: 'Full Song' })).toBeDisabled(); + expect(screen.getByRole('button', { name: 'Remix' })).toBeDisabled(); + expect(screen.getByRole('button', { name: 'Repaint' })).toBeDisabled(); + expect(screen.getByRole('button', { name: 'Separator' })).not.toBeDisabled(); + expect(screen.getByRole('button', { name: 'Download Model' })).toBeInTheDocument(); + }); + + it('shows the single local separator model and advanced settings when the model is cached', async () => { + localModelCached = true; + selectedRegionIds = ['audio-region-1']; + + render(); + + await screen.findByText('Selected Region'); + const options = await screen.findAllByRole('option'); + expect(options).toHaveLength(1); + expect(options[0]).toHaveTextContent('Vocal and Instrument (Medium Accuracy)'); + + fireEvent.click(screen.getByRole('button', { name: /Advanced Settings/i })); + expect(screen.getByLabelText('Optional audio chunk duration (seconds)')).toBeInTheDocument(); + expect(screen.getByLabelText('MDX overlap')).toBeInTheDocument(); + }); + + it('prompts for an audio region when the model is cached but nothing is selected', async () => { + localModelCached = true; + + render(); + + expect(await screen.findByText(/Select an audio region on the timeline/)).toBeInTheDocument(); + }); + + it('renders local separation outputs after processing completes', async () => { + localModelCached = true; + selectedRegionIds = ['audio-region-1']; + + render(); + + fireEvent.click(await screen.findByRole('button', { name: 'Separate Stems' })); + + await waitFor(() => { + expect(screen.getByText('Instrumental')).toBeInTheDocument(); + expect(screen.getByText('Vocals')).toBeInTheDocument(); + expect(screen.getByRole('button', { name: 'Import All Stems to Timeline' })).toBeInTheDocument(); + }); + }); +}); diff --git a/src/components/KGOnePanel.tsx b/src/components/KGOnePanel.tsx index d53dbd2..26012dd 100644 --- a/src/components/KGOnePanel.tsx +++ b/src/components/KGOnePanel.tsx @@ -15,18 +15,31 @@ import type { KeySignature } from '../core/KGProject'; import { ImportStemsCommand } from '../core/commands'; import type { StemImportEntry } from '../core/commands'; import { showAlert } from '../util/dialogUtil'; +import { + LOCAL_SEPARATOR_MODEL_CONFIG, + LOCAL_SEPARATOR_MODEL_FILENAME, + LOCAL_SEPARATOR_MODEL_URL, +} from '../util/localSeparatorConfig'; +import { LocalSeparatorModelCache } from '../util/localSeparatorModelCache'; +import { runLocalSeparator } from '../util/localSeparatorRunner'; +import { LocalOrtRuntimeManager, detectLocalRuntimeSupport } from '../util/localSeparatorRuntime'; // ─── Types ──────────────────────────────────────────────────────────────────── type Tab = 'clip' | 'fullsong' | 'remix' | 'repaint' | 'separator'; +type KGOneMode = 'server' | 'local-separator'; type GenStatus = 'idle' | 'loading-model' | 'generating' | 'polling' | 'downloading' | 'done' | 'error'; -const SEPARATOR_MODELS = [ +const SERVER_SEPARATOR_MODELS = [ { label: 'Vocal and Instrument (Medium Accuracy)', value: 'UVR-MDX-NET-Inst_HQ_3.onnx' }, { label: 'Vocal and Instrument (High Accuracy)', value: 'MDX23C-8KFFT-InstVoc_HQ.ckpt' }, { label: 'Vocal, Drums, Bass, Guitar, Piano, and Others', value: 'htdemucs_6s.yaml' }, ] as const; +const LOCAL_SEPARATOR_MODELS = [ + { label: LOCAL_SEPARATOR_MODEL_CONFIG.displayName, value: LOCAL_SEPARATOR_MODEL_FILENAME }, +] as const; +const KGONE_TABS = ['fullsong', 'remix', 'repaint', 'separator'] as const; const CLIP_NOTES = ['C', 'C#', 'D', 'D#', 'E', 'F', 'F#', 'G', 'G#', 'A', 'A#', 'B']; @@ -59,6 +72,22 @@ function formatTime(sec: number): string { return `${m}:${s.toString().padStart(2, '0')}`; } +function formatKGOneTabLabel(tab: Tab): string { + if (tab === 'fullsong') return 'Full Song'; + if (tab === 'remix') return 'Remix'; + if (tab === 'repaint') return 'Repaint'; + return 'Separator'; +} + +export function getDefaultKGOneTab(mode: KGOneMode): Tab { + return mode === 'local-separator' ? 'separator' : 'fullsong'; +} + +export function getKGOneMode(): KGOneMode { + const enabled = (ConfigManager.instance().get('general.kgone.enabled') as boolean | undefined) ?? false; + return enabled ? 'server' : 'local-separator'; +} + // ─── Shared components ──────────────────────────────────────────────────────── interface ExpanderProps { @@ -830,18 +859,33 @@ function countRepaintTracks(sourceTrackName: string): number { return tracks.filter(t => pattern.test(t.getName())).length; } -const SeparatorTab: React.FC = () => { +const SeparatorTab: React.FC<{ mode: KGOneMode }> = ({ mode }) => { const { selectedRegionIds, projectName, bpm, timeSignature, maxBars, refreshProjectState } = useProjectStore(); - const [model, setModel] = useState(SEPARATOR_MODELS[0].value); + const localOnlyMode = mode === 'local-separator'; + const availableSeparatorModels = localOnlyMode ? LOCAL_SEPARATOR_MODELS : SERVER_SEPARATOR_MODELS; + const [model, setModel] = useState(availableSeparatorModels[0].value); // Generation state const [genStatus, setGenStatus] = useState('idle'); const [genHint, setGenHint] = useState(''); const [errorMsg, setErrorMsg] = useState(''); const [stemAudioUrls, setStemAudioUrls] = useState>([]); + const runtimeSupport = useMemo(() => detectLocalRuntimeSupport(), []); + const [localProviderLabel, setLocalProviderLabel] = useState( + runtimeSupport.webgpuExposed ? 'webgpu available' : 'cpu/wasm only', + ); + const [isLocalModelCached, setIsLocalModelCached] = useState(false); + const [isCheckingLocalModel, setIsCheckingLocalModel] = useState(false); + const [isDownloadingLocalModel, setIsDownloadingLocalModel] = useState(false); + const [isDeletingLocalModel, setIsDeletingLocalModel] = useState(false); + const [localProgressPercent, setLocalProgressPercent] = useState(0); + const [localProgressText, setLocalProgressText] = useState(''); + const [localChunkDurationSeconds, setLocalChunkDurationSeconds] = useState(''); + const [localOverlap, setLocalOverlap] = useState(String(LOCAL_SEPARATOR_MODEL_CONFIG.defaults.overlap)); const abortRef = useRef(null); const taskIdRef = useRef(''); + const localRuntimeManagerRef = useRef(null); const originalRegionRef = useRef<{ regionName: string; trackName: string; @@ -861,6 +905,31 @@ const SeparatorTab: React.FC = () => { // eslint-disable-next-line react-hooks/exhaustive-deps }, []); + useEffect(() => { + setModel(availableSeparatorModels[0].value); + }, [availableSeparatorModels]); + + const refreshLocalModelCacheState = useCallback(async () => { + if (!localOnlyMode) return; + setIsCheckingLocalModel(true); + try { + setIsLocalModelCached(await LocalSeparatorModelCache.exists()); + } catch (err) { + console.error('[KGOne] Local model cache check failed:', err); + setErrorMsg(err instanceof Error ? err.message : String(err)); + } finally { + setIsCheckingLocalModel(false); + } + }, [localOnlyMode]); + + useEffect(() => { + if (!localOnlyMode) return; + localRuntimeManagerRef.current = new LocalOrtRuntimeManager({ + onProviderChange: provider => setLocalProviderLabel(provider), + }); + void refreshLocalModelCacheState(); + }, [localOnlyMode, refreshLocalModelCacheState]); + const selectedAudioRegion = useMemo(() => { if (!selectedRegionIds.length) return null; const project = KGCore.instance().getCurrentProject(); @@ -879,7 +948,57 @@ const SeparatorTab: React.FC = () => { const isGenerating = genStatus !== 'idle' && genStatus !== 'done' && genStatus !== 'error'; - const handleSeparate = useCallback(async () => { + const handleDownloadLocalModel = useCallback(async () => { + setIsDownloadingLocalModel(true); + setErrorMsg(''); + setLocalProgressPercent(0); + setLocalProgressText('Downloading local separator model...'); + + try { + await LocalSeparatorModelCache.download( + LOCAL_SEPARATOR_MODEL_URL, + LOCAL_SEPARATOR_MODEL_FILENAME, + progress => { + const receivedMb = (progress.receivedBytes / (1024 * 1024)).toFixed(1); + const totalMb = progress.totalBytes ? (progress.totalBytes / (1024 * 1024)).toFixed(1) : null; + setLocalProgressPercent(progress.totalBytes ? progress.percent : 0); + setLocalProgressText( + totalMb + ? `Downloading local separator model... ${receivedMb} / ${totalMb} MB` + : `Downloading local separator model... ${receivedMb} MB`, + ); + }, + ); + setLocalProgressPercent(100); + setLocalProgressText('Local separator model is ready.'); + await refreshLocalModelCacheState(); + } catch (err) { + setLocalProgressPercent(0); + setLocalProgressText(''); + setErrorMsg(err instanceof Error ? err.message : String(err)); + } finally { + setIsDownloadingLocalModel(false); + } + }, [refreshLocalModelCacheState]); + + const handleDeleteLocalModel = useCallback(async () => { + setIsDeletingLocalModel(true); + setErrorMsg(''); + try { + await LocalSeparatorModelCache.delete(); + localRuntimeManagerRef.current?.reset(); + setLocalProviderLabel(runtimeSupport.webgpuExposed ? 'webgpu available' : 'cpu/wasm only'); + setLocalProgressPercent(0); + setLocalProgressText(''); + await refreshLocalModelCacheState(); + } catch (err) { + setErrorMsg(err instanceof Error ? err.message : String(err)); + } finally { + setIsDeletingLocalModel(false); + } + }, [refreshLocalModelCacheState, runtimeSupport.webgpuExposed]); + + const handleSeparateServer = useCallback(async () => { if (!selectedAudioRegion) return; // Capture snapshot before anything changes — selection may shift during generation @@ -1052,7 +1171,107 @@ const SeparatorTab: React.FC = () => { setGenStatus('error'); setGenHint(''); } - }, [selectedAudioRegion, projectName, model, stemAudioUrls]); + }, [selectedAudioRegion, projectName, model, stemAudioUrls, bpm]); + + const handleSeparateLocal = useCallback(async () => { + if (!selectedAudioRegion || !isLocalModelCached) return; + + originalRegionRef.current = { + regionName: selectedAudioRegion.region.getName(), + trackName: selectedAudioRegion.trackName, + startFromBeat: selectedAudioRegion.region.getStartFromBeat(), + trackIndex: selectedAudioRegion.trackIndex, + }; + setImportError(''); + stemAudioUrls.forEach(s => URL.revokeObjectURL(s.url)); + setStemAudioUrls([]); + setErrorMsg(''); + setGenStatus('loading-model'); + setLocalProgressPercent(0); + setLocalProgressText('Preparing ONNX Runtime session...'); + setLocalProviderLabel(runtimeSupport.webgpuExposed ? 'webgpu available' : 'cpu/wasm only'); + + try { + const modelBuffer = await LocalSeparatorModelCache.getArrayBuffer(); + const runtimeManager = localRuntimeManagerRef.current ?? new LocalOrtRuntimeManager({ + onProviderChange: provider => setLocalProviderLabel(provider), + }); + localRuntimeManagerRef.current = runtimeManager; + const runtime = await runtimeManager.ensureRuntime(LOCAL_SEPARATOR_MODEL_CONFIG, new Uint8Array(modelBuffer)); + + setGenStatus('generating'); + setLocalProgressPercent(3); + setLocalProgressText('Reading audio file...'); + + const audioFileId = selectedAudioRegion.region.getAudioFileId(); + const clipStart = selectedAudioRegion.region.getClipStartOffsetSeconds(); + const fullDuration = selectedAudioRegion.region.getAudioDurationSeconds(); + const regionLengthSec = selectedAudioRegion.region.getLength() * (60 / bpm); + const effectiveDuration = Math.min(regionLengthSec, fullDuration - clipStart); + + const rawBuffer = await KGAudioFileStorage.loadAudioFile(projectName, audioFileId); + const needsSlice = clipStart > 0.01 || effectiveDuration < fullDuration - 0.01; + const inputBuffer = needsSlice + ? await sliceAudioToWav(rawBuffer, clipStart, effectiveDuration) + : rawBuffer; + + setGenStatus('polling'); + setLocalProgressPercent(5); + setLocalProgressText('Running browser separation...'); + + const chunkDuration = localChunkDurationSeconds.trim() + ? Number.parseFloat(localChunkDurationSeconds) + : null; + const overlapValue = Number.parseFloat(localOverlap); + + const result = await runLocalSeparator({ + session: runtime.session, + runtimeProvider: runtime.provider, + modelConfig: LOCAL_SEPARATOR_MODEL_CONFIG, + audioBuffer: inputBuffer, + chunkDurationSeconds: Number.isFinite(chunkDuration) && (chunkDuration ?? 0) > 0 ? chunkDuration : null, + overlap: Number.isFinite(overlapValue) ? overlapValue : LOCAL_SEPARATOR_MODEL_CONFIG.defaults.overlap, + onProviderChange: provider => setLocalProviderLabel(provider), + onProgress: progress => { + setLocalProgressPercent(progress.percent); + const chunkSuffix = progress.totalChunks ? ` (${progress.processedChunks}/${progress.totalChunks} chunks)` : ''; + setLocalProgressText(`${progress.passLabel}${chunkSuffix}`); + }, + }); + + taskIdRef.current = `local_${Date.now()}`; + const nextStemAudioUrls = result.stems.map(stem => ({ + name: stem.name, + url: URL.createObjectURL(stem.blob), + })); + setStemAudioUrls(nextStemAudioUrls); + setGenStatus('done'); + setLocalProgressPercent(100); + setLocalProgressText('Separation complete.'); + kgoneLog('RES', 'local separator summary', result.debugSummary); + } catch (err) { + console.error('[KGOne] Local separator error:', err); + setGenStatus('error'); + setErrorMsg(err instanceof Error ? err.message : String(err)); + } + }, [ + selectedAudioRegion, + isLocalModelCached, + stemAudioUrls, + runtimeSupport.webgpuExposed, + bpm, + projectName, + localChunkDurationSeconds, + localOverlap, + ]); + + const handleSeparate = useCallback(async () => { + if (localOnlyMode) { + await handleSeparateLocal(); + return; + } + await handleSeparateServer(); + }, [handleSeparateLocal, handleSeparateServer, localOnlyMode]); const handleImportAll = useCallback(async () => { const snap = originalRegionRef.current; @@ -1068,9 +1287,9 @@ const SeparatorTab: React.FC = () => { const stems: StemImportEntry[] = await Promise.all( stemAudioUrls.map(async (stem) => { const blob = await fetch(stem.url).then(r => r.blob()); - const fileName = `KGOne_Stem_${stem.name}_${taskIdRef.current}.mp3`; + const fileName = `KGOne_Stem_${stem.name}_${taskIdRef.current}.${localOnlyMode ? 'wav' : 'mp3'}`; const fileId = `kgone_stem_${Date.now()}_${Math.random().toString(36).slice(2, 7)}`; - const audioFile = new File([blob], fileName, { type: 'audio/mpeg' }); + const audioFile = new File([blob], fileName, { type: localOnlyMode ? 'audio/wav' : 'audio/mpeg' }); const arrayBuffer = await blob.arrayBuffer(); const toneBuffer = new Tone.ToneAudioBuffer(); @@ -1114,9 +1333,19 @@ const SeparatorTab: React.FC = () => { } finally { setIsImporting(false); } - }, [stemAudioUrls, projectName, maxBars, refreshProjectState]); + }, [stemAudioUrls, projectName, maxBars, refreshProjectState, localOnlyMode]); const btnLabel = () => { + if (localOnlyMode) { + switch (genStatus) { + case 'loading-model': return 'Preparing local model...'; + case 'generating': return 'Preparing audio...'; + case 'polling': return 'Separating locally...'; + case 'downloading': return 'Finalizing...'; + default: return 'Separate Stems'; + } + } + switch (genStatus) { case 'loading-model': return 'Loading model...'; case 'generating': return 'Preparing upload...'; @@ -1128,6 +1357,68 @@ const SeparatorTab: React.FC = () => { return ( <> + {localOnlyMode && ( +
+
Local Separator Mode
+
+ Only Vocal and Instrument (Medium Accuracy) is available while not integrated with K.G.One Music Studio server. + Processing in local may take long time depending on your hardware. When fallback to CPU happens, the webpage may + temporarily hang with little or no UI response until processing advances. +
+
+
Provider: {localProviderLabel}
+
Model: {isLocalModelCached ? 'downloaded' : 'not downloaded'}
+
+ {(localProgressText || isCheckingLocalModel) && ( +
+
+
+
+
+ {isCheckingLocalModel ? 'Checking local model cache...' : localProgressText} +
+
+ )} +
+ {!isLocalModelCached ? ( + + ) : ( + <> + + + + )} +
+
+ )} + {selectedAudioRegion ? ( <>
@@ -1139,14 +1430,44 @@ const SeparatorTab: React.FC = () => {
- setModel(e.target.value as typeof SERVER_SEPARATOR_MODELS[number]['value'])}> + {availableSeparatorModels.map(m => ( ))}
+ {localOnlyMode && ( + +
+ + setLocalChunkDurationSeconds(e.target.value)} + placeholder="Leave blank to process the full region" + /> +
+
+ + setLocalOverlap(e.target.value)} + /> +
+
+ )} + {/* Stem audio players — shown once separation is complete */} {stemAudioUrls.length > 0 && (
@@ -1156,7 +1477,7 @@ const SeparatorTab: React.FC = () => {
@@ -1195,7 +1516,7 @@ const SeparatorTab: React.FC = () => { {/* Status hint below button */} - {genHint &&
{genHint}
} + {(localOnlyMode ? localProgressText : genHint) && ( +
{localOnlyMode ? localProgressText : genHint}
+ )} ) : (
- Select an audio region on the timeline to extract stems from it. - Only audio regions are supported — MIDI regions cannot be separated. + {localOnlyMode && !isLocalModelCached + ? 'Download the local separator model, then select an audio region on the timeline to extract stems from it.' + : 'Select an audio region on the timeline to extract stems from it. Only audio regions are supported — MIDI regions cannot be separated.'}
)} -
- Powered by UVR5 CLI -
+ {!localOnlyMode && ( +
+ Powered by UVR5 CLI +
+ )} ); }; @@ -2228,8 +2554,21 @@ interface KGOnePanelProps { } const KGOnePanel: React.FC = ({ isVisible }) => { - const [activeTab, setActiveTab] = useState('fullsong'); + const mode = getKGOneMode(); + const [activeTab, setActiveTab] = useState(getDefaultKGOneTab(mode)); const { bpm, keySignature } = useProjectStore(); + const disabledTabs = mode === 'local-separator' + ? new Set(['fullsong', 'remix', 'repaint']) + : new Set(); + + useEffect(() => { + setActiveTab(current => { + if (disabledTabs.has(current)) { + return 'separator'; + } + return getDefaultKGOneTab(mode) === 'separator' ? 'separator' : current; + }); + }, [mode]); return (
@@ -2239,16 +2578,17 @@ const KGOnePanel: React.FC = ({ isVisible }) => {
{/* Clip tab temporarily disabled, will enable in the future */} - {(['fullsong', 'remix', 'repaint', 'separator'] as const).map(tab => ( + {KGONE_TABS.map(tab => ( ))}
@@ -2259,7 +2599,7 @@ const KGOnePanel: React.FC = ({ isVisible }) => { {activeTab === 'fullsong' && } {activeTab === 'remix' && } {activeTab === 'repaint' && } - {activeTab === 'separator' && } + {activeTab === 'separator' && }
); diff --git a/src/components/Toolbar.tsx b/src/components/Toolbar.tsx index af37e84..b4279be 100644 --- a/src/components/Toolbar.tsx +++ b/src/components/Toolbar.tsx @@ -961,8 +961,6 @@ const Toolbar: React.FC = () => { }; // K.G.One panel toggle - const isKGOneEnabled = ConfigManager.instance().get('general.kgone.enabled') as boolean ?? false; - const handleKGOneClick = () => { if (DEBUG_MODE.TOOLBAR) { console.log("K.G.One button clicked"); @@ -1241,11 +1239,9 @@ const Toolbar: React.FC = () => {
diff --git a/src/core/io/LocalSeparatorModelCache.test.ts b/src/core/io/LocalSeparatorModelCache.test.ts new file mode 100644 index 0000000..0642e25 --- /dev/null +++ b/src/core/io/LocalSeparatorModelCache.test.ts @@ -0,0 +1,151 @@ +import { beforeEach, describe, expect, it, vi } from 'vitest'; +import { LocalSeparatorModelCache } from '../../util/localSeparatorModelCache'; + +class MockWritableFileStream { + private readonly handle: MockFileSystemFileHandle; + private chunks: Uint8Array[] = []; + + constructor(handle: MockFileSystemFileHandle) { + this.handle = handle; + } + + async write(content: ArrayBuffer | ArrayBufferView): Promise { + const bytes = content instanceof ArrayBuffer + ? new Uint8Array(content) + : new Uint8Array(content.buffer, content.byteOffset, content.byteLength); + this.chunks.push(new Uint8Array(bytes)); + } + + async close(): Promise { + const total = this.chunks.reduce((sum, chunk) => sum + chunk.byteLength, 0); + const merged = new Uint8Array(total); + let offset = 0; + for (const chunk of this.chunks) { + merged.set(chunk, offset); + offset += chunk.byteLength; + } + this.handle.setContent(merged); + } + + async abort(): Promise { + this.chunks = []; + } +} + +class MockFileSystemFileHandle { + kind = 'file' as const; + private content = new Uint8Array(); + + constructor(public readonly name: string) {} + + setContent(content: Uint8Array): void { + this.content = content; + } + + async getFile(): Promise { + return { + arrayBuffer: async () => this.content.buffer.slice(0), + } as unknown as File; + } + + async createWritable(): Promise { + return new MockWritableFileStream(this); + } +} + +class MockFileSystemDirectoryHandle { + kind = 'directory' as const; + private entries = new Map(); + + constructor(public readonly name: string) {} + + async getDirectoryHandle(name: string, options?: { create?: boolean }): Promise { + let entry = this.entries.get(name); + if (!entry || entry.kind !== 'directory') { + if (!options?.create) { + throw new DOMException(`Directory "${name}" not found`, 'NotFoundError'); + } + entry = new MockFileSystemDirectoryHandle(name); + this.entries.set(name, entry); + } + return entry as MockFileSystemDirectoryHandle; + } + + async getFileHandle(name: string, options?: { create?: boolean }): Promise { + let entry = this.entries.get(name); + if (!entry || entry.kind !== 'file') { + if (!options?.create) { + throw new DOMException(`File "${name}" not found`, 'NotFoundError'); + } + entry = new MockFileSystemFileHandle(name); + this.entries.set(name, entry); + } + return entry as MockFileSystemFileHandle; + } + + async removeEntry(name: string): Promise { + if (!this.entries.has(name)) { + throw new DOMException(`Entry "${name}" not found`, 'NotFoundError'); + } + this.entries.delete(name); + } + + clear(): void { + this.entries.clear(); + } +} + +const mockRoot = new MockFileSystemDirectoryHandle('root'); + +vi.stubGlobal('navigator', { + ...navigator, + storage: { + getDirectory: vi.fn(async () => mockRoot), + }, +}); + +describe('LocalSeparatorModelCache', () => { + beforeEach(() => { + mockRoot.clear(); + vi.restoreAllMocks(); + }); + + it('downloads and stores a model in OPFS cache', async () => { + vi.stubGlobal('fetch', vi.fn(async () => new Response(new Uint8Array([1, 2, 3]), { + status: 200, + headers: { 'Content-Length': '3' }, + }))); + + await LocalSeparatorModelCache.download('https://example.com/model.onnx', 'model.onnx'); + + expect(await LocalSeparatorModelCache.exists('model.onnx')).toBe(true); + const buffer = await LocalSeparatorModelCache.getArrayBuffer('model.onnx'); + expect(Array.from(new Uint8Array(buffer))).toEqual([1, 2, 3]); + }); + + it('replaces a broken cached file on redownload', async () => { + vi.stubGlobal('fetch', vi.fn(async () => new Response(new Uint8Array([1]), { + status: 200, + headers: { 'Content-Length': '1' }, + }))); + await LocalSeparatorModelCache.download('https://example.com/model.onnx', 'model.onnx'); + + vi.stubGlobal('fetch', vi.fn(async () => new Response(new Uint8Array([9, 8, 7, 6]), { + status: 200, + headers: { 'Content-Length': '4' }, + }))); + await LocalSeparatorModelCache.download('https://example.com/model.onnx', 'model.onnx'); + + const buffer = await LocalSeparatorModelCache.getArrayBuffer('model.onnx'); + expect(Array.from(new Uint8Array(buffer))).toEqual([9, 8, 7, 6]); + }); + + it('deletes the cached model file', async () => { + vi.stubGlobal('fetch', vi.fn(async () => new Response(new Uint8Array([1, 2]), { status: 200 }))); + await LocalSeparatorModelCache.download('https://example.com/model.onnx', 'model.onnx'); + + await LocalSeparatorModelCache.delete('model.onnx'); + + expect(await LocalSeparatorModelCache.exists('model.onnx')).toBe(false); + }); +}); diff --git a/src/util/localSeparatorConfig.ts b/src/util/localSeparatorConfig.ts new file mode 100644 index 0000000..752f468 --- /dev/null +++ b/src/util/localSeparatorConfig.ts @@ -0,0 +1,31 @@ +import type { LocalSeparatorModelConfig } from './localSeparatorTypes'; + +export const LOCAL_SEPARATOR_MODEL_URL = + 'https://huggingface.co/notabilia/uvr5-models/resolve/main/UVR-MDX-NET-Inst_HQ_3.onnx'; + +export const LOCAL_SEPARATOR_MODEL_FILENAME = 'UVR-MDX-NET-Inst_HQ_3.onnx'; + +export const LOCAL_SEPARATOR_MODEL_CONFIG: LocalSeparatorModelConfig = { + filename: LOCAL_SEPARATOR_MODEL_FILENAME, + displayName: 'Vocal and Instrument (Medium Accuracy)', + status: 'ready', + defaults: { + sampleRate: 44100, + hopLength: 1024, + segmentSize: 256, + overlap: 0.25, + batchSize: 1, + enableDenoise: false, + invertUsingSpec: false, + normalizationThreshold: 0.9, + amplificationThreshold: 0, + matchMixOverlap: 0.02, + }, + metadata: { + compensate: 1.021, + mdx_dim_f_set: 3072, + mdx_dim_t_set: 8, + mdx_n_fft_scale_set: 7680, + primary_stem: 'Instrumental', + }, +}; diff --git a/src/util/localSeparatorCpuDsp.ts b/src/util/localSeparatorCpuDsp.ts new file mode 100644 index 0000000..cda0ac3 --- /dev/null +++ b/src/util/localSeparatorCpuDsp.ts @@ -0,0 +1,139 @@ +import type { LocalSeparatorModelConfig, StereoChannels } from './localSeparatorTypes'; +import { FFT, createWindowCache, getHannPeriodic, index4d, reflectPad } from './localSeparatorShared'; + +interface SpectrogramPayload { + data: Float32Array; + dims: number[]; + frames?: number; +} + +export class LocalSeparatorCpuDsp { + public readonly window: Float32Array; + public readonly dimF: number; + public readonly forwardFft: FFT; + + private readonly nFft: number; + private readonly hopLength: number; + private readonly trim: number; + private readonly numFreqBins: number; + private readonly inverseFft: FFT; + + constructor(config: LocalSeparatorModelConfig) { + this.nFft = config.metadata.mdx_n_fft_scale_set; + this.hopLength = config.defaults.hopLength; + this.dimF = config.metadata.mdx_dim_f_set; + this.trim = Math.floor(this.nFft / 2); + this.numFreqBins = Math.floor(this.nFft / 2) + 1; + const windowCache = createWindowCache(); + this.window = getHannPeriodic(this.nFft, windowCache); + this.forwardFft = new FFT(this.nFft); + this.inverseFft = new FFT(this.nFft); + } + + public async forwardStereo(leftChunk: Float32Array, rightChunk: Float32Array): Promise { + const paddedLeft = reflectPad(leftChunk, this.trim, this.trim); + const paddedRight = reflectPad(rightChunk, this.trim, this.trim); + const frames = Math.floor((paddedLeft.length - this.nFft) / this.hopLength) + 1; + const tensor = new Float32Array(4 * this.dimF * frames); + const dims = [1, 4, this.dimF, frames]; + + for (let frameIndex = 0; frameIndex < frames; frameIndex += 1) { + const offset = frameIndex * this.hopLength; + const leftSpectrum = this.frameSpectrum(paddedLeft, offset); + const rightSpectrum = this.frameSpectrum(paddedRight, offset); + + for (let freq = 0; freq < this.dimF; freq += 1) { + tensor[index4d(dims, 0, 0, freq, frameIndex)] = freq < 3 ? 0 : leftSpectrum.real[freq]; + tensor[index4d(dims, 0, 1, freq, frameIndex)] = freq < 3 ? 0 : leftSpectrum.imag[freq]; + tensor[index4d(dims, 0, 2, freq, frameIndex)] = freq < 3 ? 0 : rightSpectrum.real[freq]; + tensor[index4d(dims, 0, 3, freq, frameIndex)] = freq < 3 ? 0 : rightSpectrum.imag[freq]; + } + } + + return { data: tensor, dims, frames }; + } + + public async inverseStereo(spectrogramPayload: SpectrogramPayload): Promise { + const spectrogram = spectrogramPayload.data; + const dims = spectrogramPayload.dims; + const [, channels, freqBins, frames] = dims; + if (channels !== 4) { + throw new Error(`Expected 4 channels in MDX spectrogram, got ${channels}`); + } + + const outputLength = ((frames - 1) * this.hopLength) + this.nFft; + const left = new Float64Array(outputLength); + const right = new Float64Array(outputLength); + const leftWindowSums = new Float64Array(outputLength); + const rightWindowSums = new Float64Array(outputLength); + + for (let frameIndex = 0; frameIndex < frames; frameIndex += 1) { + const leftFrame = this.istftFrame(spectrogram, dims, 0, 1, frameIndex, freqBins); + const rightFrame = this.istftFrame(spectrogram, dims, 2, 3, frameIndex, freqBins); + const frameOffset = frameIndex * this.hopLength; + + for (let i = 0; i < this.nFft; i += 1) { + const weightedLeft = leftFrame[i] * this.window[i]; + const weightedRight = rightFrame[i] * this.window[i]; + left[frameOffset + i] += weightedLeft; + right[frameOffset + i] += weightedRight; + const weight = this.window[i] * this.window[i]; + leftWindowSums[frameOffset + i] += weight; + rightWindowSums[frameOffset + i] += weight; + } + } + + const normalizedLeft = new Float32Array(outputLength - (this.trim * 2)); + const normalizedRight = new Float32Array(outputLength - (this.trim * 2)); + for (let i = this.trim; i < outputLength - this.trim; i += 1) { + const outIndex = i - this.trim; + normalizedLeft[outIndex] = leftWindowSums[i] > 1e-8 ? left[i] / leftWindowSums[i] : 0; + normalizedRight[outIndex] = rightWindowSums[i] > 1e-8 ? right[i] / rightWindowSums[i] : 0; + } + + return [normalizedLeft, normalizedRight]; + } + + public dispose(): void {} + + private frameSpectrum(signal: Float32Array, offset: number): { real: Float64Array; imag: Float64Array } { + const real = new Float64Array(this.nFft); + const imag = new Float64Array(this.nFft); + for (let i = 0; i < this.nFft; i += 1) { + real[i] = signal[offset + i] * this.window[i]; + } + this.forwardFft.transform(real, imag); + return { real, imag }; + } + + private istftFrame( + spectrogram: Float32Array, + dims: number[], + realChannel: number, + imagChannel: number, + frameIndex: number, + freqBins: number, + ): Float32Array { + const real = new Float64Array(this.nFft); + const imag = new Float64Array(this.nFft); + + for (let freq = 0; freq < freqBins; freq += 1) { + real[freq] = spectrogram[index4d(dims, 0, realChannel, freq, frameIndex)]; + imag[freq] = spectrogram[index4d(dims, 0, imagChannel, freq, frameIndex)]; + } + + for (let freq = 1; freq < this.numFreqBins - 1; freq += 1) { + const mirrored = this.nFft - freq; + real[mirrored] = real[freq]; + imag[mirrored] = -imag[freq]; + } + + this.inverseFft.inverse(real, imag); + + const frame = new Float32Array(this.nFft); + for (let i = 0; i < this.nFft; i += 1) { + frame[i] = real[i]; + } + return frame; + } +} diff --git a/src/util/localSeparatorGpuDsp.ts b/src/util/localSeparatorGpuDsp.ts new file mode 100644 index 0000000..dad4ba0 --- /dev/null +++ b/src/util/localSeparatorGpuDsp.ts @@ -0,0 +1,220 @@ +import type { LocalSeparatorModelConfig } from './localSeparatorTypes'; +import { LocalSeparatorCpuDsp } from './localSeparatorCpuDsp'; +import { reflectPad } from './localSeparatorShared'; + +type GPUDeviceLike = any; +type GPUBufferLike = any; +type GPUComputePipelineLike = any; + +declare const GPUBufferUsage: any; +declare const GPUMapMode: any; + +const FRAMING_SHADER = ` +struct Params { + nfft: u32, + hop: u32, + frames: u32, + paddedLength: u32, +} + +@group(0) @binding(0) var leftInput: array; +@group(0) @binding(1) var rightInput: array; +@group(0) @binding(2) var window: array; +@group(0) @binding(3) var output: array; +@group(0) @binding(4) var params: Params; + +@compute @workgroup_size(256) +fn main(@builtin(global_invocation_id) gid: vec3) { + let index = gid.x; + let total = params.frames * params.nfft * 2u; + if (index >= total) { + return; + } + + let sample = index % params.nfft; + let frame = (index / params.nfft) % params.frames; + let channel = index / (params.nfft * params.frames); + let sourceIndex = frame * params.hop + sample; + let sampleValue = select(leftInput[sourceIndex], rightInput[sourceIndex], channel == 1u); + output[index] = sampleValue * window[sample]; +} +`; + +function alignTo(value: number, alignment: number): number { + return Math.ceil(value / alignment) * alignment; +} + +async function readBuffer(device: GPUDeviceLike, sourceBuffer: GPUBufferLike, size: number): Promise { + const readBuffer = device.createBuffer({ + size: alignTo(size, 4), + usage: GPUBufferUsage.COPY_DST | GPUBufferUsage.MAP_READ, + }); + + const encoder = device.createCommandEncoder(); + encoder.copyBufferToBuffer(sourceBuffer, 0, readBuffer, 0, size); + device.queue.submit([encoder.finish()]); + + await readBuffer.mapAsync(GPUMapMode.READ); + const copy = new Float32Array(readBuffer.getMappedRange().slice(0)); + readBuffer.unmap(); + readBuffer.destroy(); + return copy; +} + +export class LocalSeparatorGpuDsp { + private readonly device: GPUDeviceLike; + private readonly cpuDsp: LocalSeparatorCpuDsp; + private readonly nFft: number; + private readonly hopLength: number; + private readonly trim: number; + private windowBuffer: GPUBufferLike | null = null; + private paramBuffer: GPUBufferLike | null = null; + private readonly pipeline: GPUComputePipelineLike; + + public static async create(config: LocalSeparatorModelConfig): Promise { + if (!('gpu' in navigator)) { + throw new Error('WebGPU is not available for GPU DSP.'); + } + + const adapter = await (navigator as { gpu?: { requestAdapter: (options: { powerPreference: string }) => Promise } }).gpu?.requestAdapter({ + powerPreference: 'high-performance', + }); + if (!adapter) { + throw new Error('No WebGPU adapter was available for GPU DSP.'); + } + + const device = await adapter.requestDevice(); + return new LocalSeparatorGpuDsp(config, device); + } + + private constructor(config: LocalSeparatorModelConfig, device: GPUDeviceLike) { + this.device = device; + this.cpuDsp = new LocalSeparatorCpuDsp(config); + this.nFft = config.metadata.mdx_n_fft_scale_set; + this.hopLength = config.defaults.hopLength; + this.trim = Math.floor(this.nFft / 2); + this.pipeline = device.createComputePipeline({ + layout: 'auto', + compute: { + module: device.createShaderModule({ code: FRAMING_SHADER }), + entryPoint: 'main', + }, + }); + } + + public async forwardStereo(leftChunk: Float32Array, rightChunk: Float32Array): Promise<{ + data: Float32Array; + dims: number[]; + frames: number; + }> { + this.ensureStaticBuffers(); + + const paddedLeft = reflectPad(leftChunk, this.trim, this.trim); + const paddedRight = reflectPad(rightChunk, this.trim, this.trim); + const frames = Math.floor((paddedLeft.length - this.nFft) / this.hopLength) + 1; + + const leftBuffer = this.device.createBuffer({ + size: paddedLeft.byteLength, + usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST, + }); + const rightBuffer = this.device.createBuffer({ + size: paddedRight.byteLength, + usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST, + }); + const framedSize = frames * this.nFft * 2 * 4; + const outputBuffer = this.device.createBuffer({ + size: framedSize, + usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC, + }); + + this.device.queue.writeBuffer(leftBuffer, 0, paddedLeft); + this.device.queue.writeBuffer(rightBuffer, 0, paddedRight); + this.device.queue.writeBuffer(this.paramBuffer!, 0, new Uint32Array([this.nFft, this.hopLength, frames, paddedLeft.length])); + + const bindGroup = this.device.createBindGroup({ + layout: this.pipeline.getBindGroupLayout(0), + entries: [ + { binding: 0, resource: { buffer: leftBuffer } }, + { binding: 1, resource: { buffer: rightBuffer } }, + { binding: 2, resource: { buffer: this.windowBuffer! } }, + { binding: 3, resource: { buffer: outputBuffer } }, + { binding: 4, resource: { buffer: this.paramBuffer! } }, + ], + }); + + const encoder = this.device.createCommandEncoder(); + const pass = encoder.beginComputePass(); + pass.setPipeline(this.pipeline); + pass.setBindGroup(0, bindGroup); + pass.dispatchWorkgroups(Math.ceil((frames * this.nFft * 2) / 256)); + pass.end(); + this.device.queue.submit([encoder.finish()]); + + const framed = await readBuffer(this.device, outputBuffer, framedSize); + + leftBuffer.destroy(); + rightBuffer.destroy(); + outputBuffer.destroy(); + + return this.packFramedAudio(framed, frames); + } + + public async inverseStereo(payload: { data: Float32Array; dims: number[]; frames?: number }) { + return this.cpuDsp.inverseStereo(payload); + } + + public dispose(): void { + this.cpuDsp.dispose(); + this.windowBuffer?.destroy(); + this.paramBuffer?.destroy(); + } + + private ensureStaticBuffers(): void { + if (!this.windowBuffer) { + const window = this.cpuDsp.window; + this.windowBuffer = this.device.createBuffer({ + size: window.byteLength, + usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST, + }); + this.device.queue.writeBuffer(this.windowBuffer, 0, window); + } + + if (!this.paramBuffer) { + this.paramBuffer = this.device.createBuffer({ + size: 16, + usage: GPUBufferUsage.UNIFORM | GPUBufferUsage.COPY_DST, + }); + } + } + + private packFramedAudio(framed: Float32Array, frames: number): { data: Float32Array; dims: number[]; frames: number } { + const tensor = new Float32Array(4 * this.cpuDsp.dimF * frames); + const dims = [1, 4, this.cpuDsp.dimF, frames]; + + for (let frameIndex = 0; frameIndex < frames; frameIndex += 1) { + const leftOffset = frameIndex * this.nFft; + const rightOffset = (frames * this.nFft) + leftOffset; + const leftSpectrum = this.fftFrame(framed, leftOffset); + const rightSpectrum = this.fftFrame(framed, rightOffset); + + for (let freq = 0; freq < this.cpuDsp.dimF; freq += 1) { + tensor[((freq * frames) + frameIndex)] = freq < 3 ? 0 : leftSpectrum.real[freq]; + tensor[(this.cpuDsp.dimF * frames) + ((freq * frames) + frameIndex)] = freq < 3 ? 0 : leftSpectrum.imag[freq]; + tensor[(2 * this.cpuDsp.dimF * frames) + ((freq * frames) + frameIndex)] = freq < 3 ? 0 : rightSpectrum.real[freq]; + tensor[(3 * this.cpuDsp.dimF * frames) + ((freq * frames) + frameIndex)] = freq < 3 ? 0 : rightSpectrum.imag[freq]; + } + } + + return { data: tensor, dims, frames }; + } + + private fftFrame(framed: Float32Array, offset: number): { real: Float64Array; imag: Float64Array } { + const real = new Float64Array(this.nFft); + const imag = new Float64Array(this.nFft); + for (let i = 0; i < this.nFft; i += 1) { + real[i] = framed[offset + i]; + } + this.cpuDsp.forwardFft.transform(real, imag); + return { real, imag }; + } +} diff --git a/src/util/localSeparatorModelCache.ts b/src/util/localSeparatorModelCache.ts new file mode 100644 index 0000000..a4990ac --- /dev/null +++ b/src/util/localSeparatorModelCache.ts @@ -0,0 +1,126 @@ +import { LOCAL_SEPARATOR_MODEL_FILENAME } from './localSeparatorConfig'; + +export interface ModelDownloadProgress { + receivedBytes: number; + totalBytes: number | null; + percent: number; +} + +export class LocalSeparatorModelCache { + private static readonly MODELS_DIR = 'models'; + private static readonly TEMP_SUFFIX = '.download'; + + public static async exists(filename: string = LOCAL_SEPARATOR_MODEL_FILENAME): Promise { + try { + const modelsDir = await this.getModelsDir(); + await modelsDir.getFileHandle(filename); + return true; + } catch { + return false; + } + } + + public static async getFile(filename: string = LOCAL_SEPARATOR_MODEL_FILENAME): Promise { + const modelsDir = await this.getModelsDir(); + const fileHandle = await modelsDir.getFileHandle(filename); + return fileHandle.getFile(); + } + + public static async getArrayBuffer(filename: string = LOCAL_SEPARATOR_MODEL_FILENAME): Promise { + const file = await this.getFile(filename); + return file.arrayBuffer(); + } + + public static async delete(filename: string = LOCAL_SEPARATOR_MODEL_FILENAME): Promise { + try { + const modelsDir = await this.getModelsDir(); + await modelsDir.removeEntry(filename); + } catch { + // Ignore missing file cleanup. + } + + try { + const modelsDir = await this.getModelsDir(); + await modelsDir.removeEntry(`${filename}${this.TEMP_SUFFIX}`); + } catch { + // Ignore missing temp file cleanup. + } + } + + public static async download( + sourceUrl: string, + filename: string = LOCAL_SEPARATOR_MODEL_FILENAME, + onProgress?: (progress: ModelDownloadProgress) => void, + ): Promise { + const response = await fetch(sourceUrl); + if (!response.ok) { + throw new Error(`Model download failed (${response.status})`); + } + + const modelsDir = await this.getModelsDir(); + const tempName = `${filename}${this.TEMP_SUFFIX}`; + await this.delete(filename); + + const tempHandle = await modelsDir.getFileHandle(tempName, { create: true }); + const writable = await tempHandle.createWritable(); + + try { + const totalBytesHeader = response.headers.get('Content-Length'); + const totalBytes = totalBytesHeader ? Number(totalBytesHeader) : null; + + if (!response.body) { + const buffer = await response.arrayBuffer(); + await writable.write(buffer); + onProgress?.({ + receivedBytes: buffer.byteLength, + totalBytes, + percent: 100, + }); + } else { + const reader = response.body.getReader(); + let receivedBytes = 0; + + while (true) { + const { done, value } = await reader.read(); + if (done) break; + if (!value) continue; + await writable.write(value); + receivedBytes += value.byteLength; + onProgress?.({ + receivedBytes, + totalBytes, + percent: totalBytes ? (receivedBytes / totalBytes) * 100 : 0, + }); + } + } + } catch (error) { + await writable.abort(); + await this.delete(filename); + throw error; + } + + await writable.close(); + + const finalHandle = await modelsDir.getFileHandle(filename, { create: true }); + const finalWritable = await finalHandle.createWritable(); + try { + const tempFile = await tempHandle.getFile(); + await finalWritable.write(await tempFile.arrayBuffer()); + await finalWritable.close(); + } catch (error) { + await finalWritable.abort(); + throw error; + } finally { + try { + await modelsDir.removeEntry(tempName); + } catch { + // Ignore temp cleanup errors. + } + } + } + + private static async getModelsDir(): Promise { + const root = await navigator.storage.getDirectory(); + return root.getDirectoryHandle(this.MODELS_DIR, { create: true }); + } +} diff --git a/src/util/localSeparatorRunner.ts b/src/util/localSeparatorRunner.ts new file mode 100644 index 0000000..8419093 --- /dev/null +++ b/src/util/localSeparatorRunner.ts @@ -0,0 +1,589 @@ +import * as ort from 'onnxruntime-web/webgpu'; +import { LocalSeparatorCpuDsp } from './localSeparatorCpuDsp'; +import { LocalSeparatorGpuDsp } from './localSeparatorGpuDsp'; +import { LocalSeparatorTimingCollector } from './localSeparatorTiming'; +import type { + LocalRuntimeProvider, + LocalSeparatorModelConfig, + LocalSeparatorProgress, + StereoChannels, +} from './localSeparatorTypes'; +import { + concatFloat32, + createWindowCache, + getHanning, + negateArray, + normalizeChannels, + scaleChannels, + sliceChannels, +} from './localSeparatorShared'; + +const SAMPLE_RATE = 44100; + +function localSeparatorLog(message: string, payload?: unknown): void { + if (payload === undefined) { + console.log(`[localSeparator] ${message}`); + return; + } + console.log(`[localSeparator] ${message}`, payload); +} + +interface BrowserMdxSeparatorOptions { + overlap?: number; + runtimeBatchSize?: number; + timing?: LocalSeparatorTimingCollector; + onProgress?: (progress: LocalSeparatorProgress) => void; + onProviderChange?: (provider: string) => void; +} + +interface SpectrogramPayload { + data: Float32Array; + dims: number[]; + frames?: number; +} + +function packBatchPayloads(payloads: SpectrogramPayload[]): { data: Float32Array; dims: number[]; itemSize: number } { + const frames = payloads[0].dims[3]; + const dimF = payloads[0].dims[2]; + const batch = payloads.length; + const itemSize = 4 * dimF * frames; + const data = new Float32Array(batch * itemSize); + + payloads.forEach((payload, index) => { + data.set(payload.data, index * itemSize); + }); + + return { + data, + dims: [batch, 4, dimF, frames], + itemSize, + }; +} + +function unpackBatchOutput(outputData: Float32Array, batchInfo: { dims: number[]; itemSize: number }): SpectrogramPayload[] { + const results: SpectrogramPayload[] = []; + for (let index = 0; index < batchInfo.dims[0]; index += 1) { + const start = index * batchInfo.itemSize; + const end = start + batchInfo.itemSize; + results.push({ + data: outputData.slice(start, end), + dims: [1, 4, batchInfo.dims[2], batchInfo.dims[3]], + }); + } + return results; +} + +class BrowserMdxSeparator { + private readonly session: ort.InferenceSession; + private readonly runtimeProvider: LocalRuntimeProvider; + private readonly defaults: LocalSeparatorModelConfig['defaults']; + private readonly metadata: LocalSeparatorModelConfig['metadata']; + public onProgress: (progress: LocalSeparatorProgress) => void; + private overlap: number; + private runtimeBatchSize: number; + private readonly enableDenoise: boolean; + private readonly compensate: number; + private readonly primaryStem: string; + private readonly secondaryStem: string; + private readonly nFft: number; + private readonly hopLength: number; + private readonly chunkSize: number; + private readonly trim: number; + private readonly windowCache = createWindowCache(); + private readonly timing: LocalSeparatorTimingCollector; + private dsp: LocalSeparatorCpuDsp | LocalSeparatorGpuDsp; + private dspMode: 'cpu' | 'gpu-hybrid'; + + public static async create( + session: ort.InferenceSession, + runtimeProvider: LocalRuntimeProvider, + config: LocalSeparatorModelConfig, + options: BrowserMdxSeparatorOptions = {}, + ): Promise { + const timing = options.timing ?? new LocalSeparatorTimingCollector('mdx-separation'); + let dsp: LocalSeparatorCpuDsp | LocalSeparatorGpuDsp | null = null; + let dspMode: 'cpu' | 'gpu-hybrid' = 'cpu'; + + if (runtimeProvider === 'webgpu') { + try { + dsp = await timing.measureAsync('dspInit', () => LocalSeparatorGpuDsp.create(config)); + dspMode = 'gpu-hybrid'; + localSeparatorLog('GPU DSP initialized successfully.'); + } catch (error) { + console.warn('[localSeparator] GPU DSP initialization failed, using CPU DSP.', error); + options.onProviderChange?.('cpu/wasm fallback'); + localSeparatorLog('GPU DSP initialization failed. Falling back to CPU DSP.', error); + } + } + + if (!dsp) { + dsp = new LocalSeparatorCpuDsp(config); + localSeparatorLog('Using CPU DSP.'); + } + + return new BrowserMdxSeparator(session, runtimeProvider, config, { + ...options, + dsp, + dspMode, + timing, + }); + } + + private constructor( + session: ort.InferenceSession, + runtimeProvider: LocalRuntimeProvider, + config: LocalSeparatorModelConfig, + options: BrowserMdxSeparatorOptions & { + dsp: LocalSeparatorCpuDsp | LocalSeparatorGpuDsp; + dspMode: 'cpu' | 'gpu-hybrid'; + timing: LocalSeparatorTimingCollector; + }, + ) { + this.session = session; + this.runtimeProvider = runtimeProvider; + this.defaults = config.defaults; + this.metadata = config.metadata; + this.onProgress = options.onProgress ?? (() => {}); + this.overlap = options.overlap ?? this.defaults.overlap; + this.runtimeBatchSize = Math.max(1, options.runtimeBatchSize ?? 2); + this.enableDenoise = this.defaults.enableDenoise; + this.compensate = this.metadata.compensate; + this.primaryStem = this.metadata.primary_stem ?? 'Vocals'; + this.secondaryStem = this.primaryStem === 'Instrumental' ? 'Vocals' : 'Instrumental'; + this.nFft = this.metadata.mdx_n_fft_scale_set; + this.hopLength = this.defaults.hopLength; + this.trim = Math.floor(this.nFft / 2); + this.chunkSize = this.hopLength * (this.defaults.segmentSize - 1); + this.dsp = options.dsp; + this.dspMode = options.dspMode; + this.timing = options.timing; + } + + public dispose(): void { + this.dsp.dispose(); + } + + public getDebugSummary(extra: Record = {}): Record { + return this.timing.getSummary({ + runtimeProvider: this.runtimeProvider, + dspMode: this.dspMode, + ...extra, + }); + } + + public async separate(channels: StereoChannels): Promise<{ + stems: Record; + primaryStem: string; + secondaryStem: string; + }> { + this.onProgress({ stage: 'main', passLabel: 'Main pass', percent: 0, processedChunks: 0, totalChunks: 0 }); + + const { channels: normalizedChannels, originalPeak } = this.timing.measureSync( + 'normalize', + () => normalizeChannels(channels, this.defaults.normalizationThreshold, this.defaults.amplificationThreshold), + ); + const primarySource = await this.demix(normalizedChannels, false); + this.onProgress({ + stage: 'main-complete', + passLabel: 'Main pass', + percent: this.defaults.invertUsingSpec ? 50 : 100, + processedChunks: 0, + totalChunks: 0, + }); + + const primaryScaled = this.timing.measureSync('scalePrimary', () => scaleChannels(primarySource, originalPeak)); + + let secondaryChannels: StereoChannels; + if (this.defaults.invertUsingSpec) { + const rawMix = await this.demix(normalizedChannels, true); + secondaryChannels = this.timing.measureSync('secondaryFromMix', () => { + const secondaryLeft = new Float32Array(rawMix[0].length); + const secondaryRight = new Float32Array(rawMix[1].length); + + for (let i = 0; i < secondaryLeft.length; i += 1) { + secondaryLeft[i] = rawMix[0][i] - (primaryScaled[0][i] * this.compensate); + secondaryRight[i] = rawMix[1][i] - (primaryScaled[1][i] * this.compensate); + } + + return [secondaryLeft, secondaryRight]; + }); + } else { + secondaryChannels = this.timing.measureSync('secondarySubtract', () => { + const secondaryLeft = new Float32Array(primaryScaled[0].length); + const secondaryRight = new Float32Array(primaryScaled[1].length); + for (let i = 0; i < secondaryLeft.length; i += 1) { + secondaryLeft[i] = normalizedChannels[0][i] - (primaryScaled[0][i] * this.compensate); + secondaryRight[i] = normalizedChannels[1][i] - (primaryScaled[1][i] * this.compensate); + } + return [secondaryLeft, secondaryRight]; + }); + } + + const primaryNormalized = this.timing.measureSync( + 'normalizePrimaryOutput', + () => normalizeChannels(primaryScaled, this.defaults.normalizationThreshold, this.defaults.amplificationThreshold).channels, + ); + const secondaryNormalized = this.timing.measureSync( + 'normalizeSecondaryOutput', + () => normalizeChannels(secondaryChannels, this.defaults.normalizationThreshold, this.defaults.amplificationThreshold).channels, + ); + + return { + stems: { + [this.primaryStem]: primaryNormalized, + [this.secondaryStem]: secondaryNormalized, + }, + primaryStem: this.primaryStem, + secondaryStem: this.secondaryStem, + }; + } + + private async demix(channels: StereoChannels, isMatchMix: boolean): Promise { + const overlap = isMatchMix ? this.defaults.matchMixOverlap : this.overlap; + const genSize = this.chunkSize - (2 * this.trim); + const pad = genSize + this.trim - (channels[0].length % genSize); + const mixture: StereoChannels = [ + concatFloat32([new Float32Array(this.trim), channels[0], new Float32Array(pad)]), + concatFloat32([new Float32Array(this.trim), channels[1], new Float32Array(pad)]), + ]; + + const step = Math.max(1, Math.trunc((1 - overlap) * this.chunkSize)); + const result: StereoChannels = [new Float32Array(mixture[0].length), new Float32Array(mixture[1].length)]; + const divider: StereoChannels = [new Float32Array(mixture[0].length), new Float32Array(mixture[1].length)]; + const totalChunks = Math.ceil(mixture[0].length / step); + let processedChunks = 0; + + const windows: Array<{ + start: number; + actualSize: number; + leftChunk: Float32Array; + rightChunk: Float32Array; + window: Float32Array | null; + }> = []; + + for (let start = 0; start < mixture[0].length; start += step) { + const end = Math.min(start + this.chunkSize, mixture[0].length); + const actualSize = end - start; + const leftChunk = new Float32Array(this.chunkSize); + const rightChunk = new Float32Array(this.chunkSize); + leftChunk.set(mixture[0].subarray(start, end)); + rightChunk.set(mixture[1].subarray(start, end)); + windows.push({ + start, + actualSize, + leftChunk, + rightChunk, + window: overlap !== 0 ? getHanning(actualSize, this.windowCache) : null, + }); + } + + for (let batchStart = 0; batchStart < windows.length; batchStart += this.runtimeBatchSize) { + const batch = windows.slice(batchStart, batchStart + this.runtimeBatchSize); + const tarWavesBatch = await this.processBatch(batch, isMatchMix); + + batch.forEach((chunk, index) => { + const tarWaves = tarWavesBatch[index]; + for (let i = 0; i < chunk.actualSize; i += 1) { + const weight = chunk.window ? chunk.window[i] : 1; + result[0][chunk.start + i] += tarWaves[0][i] * weight; + result[1][chunk.start + i] += tarWaves[1][i] * weight; + divider[0][chunk.start + i] += weight; + divider[1][chunk.start + i] += weight; + } + + processedChunks += 1; + const passFraction = totalChunks > 0 ? processedChunks / totalChunks : 1; + const overallPercent = isMatchMix ? 50 + (passFraction * 50) : passFraction * (this.defaults.invertUsingSpec ? 50 : 100); + this.onProgress({ + stage: isMatchMix ? 'match-mix' : 'main', + passLabel: isMatchMix ? 'Match-mix pass' : 'Main pass', + percent: overallPercent, + processedChunks, + totalChunks, + }); + }); + } + + const left = new Float32Array(channels[0].length); + const right = new Float32Array(channels[1].length); + const endTrim = result[0].length - this.trim; + for (let i = this.trim; i < endTrim; i += 1) { + const outIndex = i - this.trim; + if (outIndex >= left.length) break; + left[outIndex] = divider[0][i] > 1e-8 ? result[0][i] / divider[0][i] : 0; + right[outIndex] = divider[1][i] > 1e-8 ? result[1][i] / divider[1][i] : 0; + } + + return [left, right]; + } + + private async processBatch( + batch: Array<{ leftChunk: Float32Array; rightChunk: Float32Array }>, + isMatchMix: boolean, + ): Promise { + const spectra = await this.timing.measureAsync( + isMatchMix ? 'matchMixDspForward' : 'dspForward', + () => Promise.all(batch.map(chunk => this.dsp.forwardStereo(chunk.leftChunk, chunk.rightChunk))), + ); + + if (isMatchMix) { + return this.timing.measureAsync( + 'matchMixDspInverse', + () => Promise.all(spectra.map(payload => this.dsp.inverseStereo(payload))), + ); + } + + let predictedPayloads; + if (this.enableDenoise) { + const positiveOutput = await this.executeModelBatch(spectra); + const negativePayloads = spectra.map(payload => ({ + data: negateArray(payload.data), + dims: payload.dims, + })); + const negativeOutput = await this.executeModelBatch(negativePayloads); + predictedPayloads = positiveOutput.map((payload, index) => { + const data = new Float32Array(payload.data.length); + for (let i = 0; i < data.length; i += 1) { + data[i] = (negativeOutput[index].data[i] * -0.5) + (payload.data[i] * 0.5); + } + return { data, dims: payload.dims }; + }); + } else { + predictedPayloads = await this.executeModelBatch(spectra); + } + + return this.timing.measureAsync( + 'dspInverse', + () => Promise.all(predictedPayloads.map(payload => this.dsp.inverseStereo(payload))), + ); + } + + private async executeModelBatch(payloads: SpectrogramPayload[]): Promise { + const packed = packBatchPayloads(payloads); + const tensor = new ort.Tensor('float32', packed.data, packed.dims); + const feeds = { [this.session.inputNames[0]]: tensor }; + try { + const outputs = await this.timing.measureAsync('inference', () => this.session.run(feeds)); + const firstOutputName = this.session.outputNames[0]; + return unpackBatchOutput(outputs[firstOutputName].data as Float32Array, packed); + } catch (error) { + if (payloads.length > 1 && this.runtimeBatchSize > 1) { + console.warn('[localSeparator] Batched inference failed, falling back to batch size 1.', error); + localSeparatorLog('Batched inference failed. Falling back to batch size 1.', error); + this.runtimeBatchSize = 1; + const singleResults: SpectrogramPayload[] = []; + for (const payload of payloads) { + const singlePacked = packBatchPayloads([payload]); + const singleTensor = new ort.Tensor('float32', singlePacked.data, singlePacked.dims); + const singleFeeds = { [this.session.inputNames[0]]: singleTensor }; + const outputs = await this.timing.measureAsync('inferenceFallback', () => this.session.run(singleFeeds)); + const firstOutputName = this.session.outputNames[0]; + singleResults.push(...unpackBatchOutput(outputs[firstOutputName].data as Float32Array, singlePacked)); + } + return singleResults; + } + throw error; + } + } +} + +export async function decodeAudioToStereo(arrayBuffer: ArrayBuffer): Promise { + const audioContext = new AudioContext({ sampleRate: SAMPLE_RATE }); + const decoded = await audioContext.decodeAudioData(arrayBuffer.slice(0)); + let buffer = decoded; + + if (decoded.sampleRate !== SAMPLE_RATE) { + const offline = new OfflineAudioContext({ + numberOfChannels: Math.max(2, decoded.numberOfChannels), + length: Math.ceil(decoded.duration * SAMPLE_RATE), + sampleRate: SAMPLE_RATE, + }); + const source = offline.createBufferSource(); + source.buffer = decoded; + source.connect(offline.destination); + source.start(); + buffer = await offline.startRendering(); + } + + await audioContext.close(); + + if (buffer.numberOfChannels === 1) { + const mono = buffer.getChannelData(0); + return [new Float32Array(mono), new Float32Array(mono)]; + } + + return [new Float32Array(buffer.getChannelData(0)), new Float32Array(buffer.getChannelData(1))]; +} + +export function channelsToWavBlob(channels: StereoChannels, sampleRate: number = SAMPLE_RATE): Blob { + const length = channels[0].length; + const interleaved = new Int16Array(length * 2); + for (let i = 0; i < length; i += 1) { + interleaved[i * 2] = toInt16(channels[0][i]); + interleaved[(i * 2) + 1] = toInt16(channels[1][i]); + } + + const buffer = new ArrayBuffer(44 + (interleaved.length * 2)); + const view = new DataView(buffer); + writeAscii(view, 0, 'RIFF'); + view.setUint32(4, 36 + (interleaved.length * 2), true); + writeAscii(view, 8, 'WAVE'); + writeAscii(view, 12, 'fmt '); + view.setUint32(16, 16, true); + view.setUint16(20, 1, true); + view.setUint16(22, 2, true); + view.setUint32(24, sampleRate, true); + view.setUint32(28, sampleRate * 4, true); + view.setUint16(32, 4, true); + view.setUint16(34, 16, true); + writeAscii(view, 36, 'data'); + view.setUint32(40, interleaved.length * 2, true); + + let offset = 44; + for (let i = 0; i < interleaved.length; i += 1) { + view.setInt16(offset, interleaved[i], true); + offset += 2; + } + + return new Blob([buffer], { type: 'audio/wav' }); +} + +function writeAscii(view: DataView, offset: number, text: string): void { + for (let i = 0; i < text.length; i += 1) { + view.setUint8(offset + i, text.charCodeAt(i)); + } +} + +function toInt16(value: number): number { + const clamped = Math.max(-1, Math.min(1, value)); + return clamped < 0 ? Math.round(clamped * 0x8000) : Math.round(clamped * 0x7fff); +} + +function concatChannelPairs(chunks: StereoChannels[]): StereoChannels { + return [concatFloat32(chunks.map(chunk => chunk[0])), concatFloat32(chunks.map(chunk => chunk[1]))]; +} + +export async function runLocalSeparator(options: { + session: ort.InferenceSession; + runtimeProvider: LocalRuntimeProvider; + modelConfig: LocalSeparatorModelConfig; + audioBuffer: ArrayBuffer; + chunkDurationSeconds: number | null; + overlap: number; + onProgress: (progress: LocalSeparatorProgress) => void; + onProviderChange?: (provider: string) => void; +}): Promise<{ + stems: Array<{ name: string; blob: Blob }>; + providerLabel: string; + debugSummary: Record; +}> { + const timing = new LocalSeparatorTimingCollector('local-separation'); + const decoded = await timing.measureAsync('decode', () => decodeAudioToStereo(options.audioBuffer)); + localSeparatorLog(`Running browser MDX separation on ${options.runtimeProvider === 'webgpu' ? 'GPU/WebGPU' : 'CPU/wasm'}...`); + + const separator = await BrowserMdxSeparator.create( + options.session, + options.runtimeProvider, + options.modelConfig, + { + overlap: options.overlap, + runtimeBatchSize: options.modelConfig.defaults.batchSize > 1 ? options.modelConfig.defaults.batchSize : 2, + timing, + onProgress: options.onProgress, + onProviderChange: options.onProviderChange, + }, + ); + + try { + const outputs = await separateWithOptionalChunking( + separator, + decoded, + timing, + options.chunkDurationSeconds, + options.modelConfig, + options.onProgress, + ); + const primaryBlob = timing.measureSync('wavEncodePrimary', () => channelsToWavBlob(outputs.stems[outputs.primaryStem])); + const secondaryBlob = timing.measureSync('wavEncodeSecondary', () => channelsToWavBlob(outputs.stems[outputs.secondaryStem])); + + return { + stems: [ + { name: outputs.primaryStem, blob: primaryBlob }, + { name: outputs.secondaryStem, blob: secondaryBlob }, + ], + providerLabel: options.runtimeProvider === 'webgpu' ? 'GPU/WebGPU' : 'CPU/wasm', + debugSummary: separator.getDebugSummary({ model: options.modelConfig.filename }), + }; + } finally { + localSeparatorLog('Separation timing summary', separator.getDebugSummary({ model: options.modelConfig.filename })); + separator.dispose(); + } +} + +async function separateWithOptionalChunking( + separator: BrowserMdxSeparator, + decoded: StereoChannels, + timing: LocalSeparatorTimingCollector, + chunkDurationSeconds: number | null, + modelConfig: LocalSeparatorModelConfig, + onProgress: (progress: LocalSeparatorProgress) => void, +): Promise<{ + stems: Record; + primaryStem: string; + secondaryStem: string; +}> { + if (!chunkDurationSeconds) { + return separator.separate(decoded); + } + + const chunkSamples = Math.max(1, Math.floor(chunkDurationSeconds * SAMPLE_RATE)); + if (decoded[0].length <= chunkSamples) { + return separator.separate(decoded); + } + + const totalChunks = Math.ceil(decoded[0].length / chunkSamples); + const primaryStem = modelConfig.metadata.primary_stem ?? 'Vocals'; + const secondaryStem = primaryStem === 'Instrumental' ? 'Vocals' : 'Instrumental'; + const primaryStemChunks: StereoChannels[] = []; + const secondaryStemChunks: StereoChannels[] = []; + const baseOnProgress = separator.onProgress; + + for (let index = 0; index < totalChunks; index += 1) { + const start = index * chunkSamples; + const end = Math.min(start + chunkSamples, decoded[0].length); + const chunk = sliceChannels(decoded, start, end); + + separator.onProgress = progress => { + const chunkFraction = progress.percent / 100; + const overallPercent = ((index + chunkFraction) / totalChunks) * 100; + baseOnProgress({ + ...progress, + percent: overallPercent, + passLabel: `Audio chunk ${index + 1}/${totalChunks}: ${progress.passLabel}`, + }); + }; + + onProgress({ + stage: 'chunk-prep', + passLabel: `Audio chunk ${index + 1}/${totalChunks}: preparing ${Math.round((end - start) / SAMPLE_RATE)}s chunk...`, + percent: (index / totalChunks) * 100, + processedChunks: index, + totalChunks, + }); + + const result = await timing.measureAsync('chunkedSeparate', () => separator.separate(chunk)); + primaryStemChunks.push(result.stems[primaryStem]); + secondaryStemChunks.push(result.stems[secondaryStem]); + } + + separator.onProgress = baseOnProgress; + + return { + stems: { + [primaryStem]: concatChannelPairs(primaryStemChunks), + [secondaryStem]: concatChannelPairs(secondaryStemChunks), + }, + primaryStem, + secondaryStem, + }; +} diff --git a/src/util/localSeparatorRuntime.ts b/src/util/localSeparatorRuntime.ts new file mode 100644 index 0000000..b44b17c --- /dev/null +++ b/src/util/localSeparatorRuntime.ts @@ -0,0 +1,95 @@ +import * as ort from 'onnxruntime-web/webgpu'; +import ortWasmJsepMjsUrl from 'onnxruntime-web/ort-wasm-simd-threaded.jsep.mjs?url'; +import ortWasmJsepUrl from 'onnxruntime-web/ort-wasm-simd-threaded.jsep.wasm?url'; +import type { LocalRuntimeState, LocalRuntimeSupport, LocalSeparatorModelConfig } from './localSeparatorTypes'; + +function localSeparatorLog(message: string, payload?: unknown): void { + if (payload === undefined) { + console.log(`[localSeparator] ${message}`); + return; + } + console.log(`[localSeparator] ${message}`, payload); +} + +export function detectLocalRuntimeSupport(): LocalRuntimeSupport { + const support = { + webgpuExposed: typeof navigator !== 'undefined' && 'gpu' in navigator, + }; + if (support.webgpuExposed) { + localSeparatorLog('WebGPU API is exposed by this browser.'); + } else { + localSeparatorLog('WebGPU API is not exposed by this browser. CPU/wasm will be used.'); + } + return support; +} + +export class LocalOrtRuntimeManager { + private runtime: LocalRuntimeState | null = null; + private currentModel: string | null = null; + private readonly onProviderChange: (provider: string) => void; + private static wasmPathsConfigured = false; + + constructor({ onProviderChange }: { onProviderChange?: (provider: string) => void } = {}) { + this.onProviderChange = onProviderChange ?? (() => {}); + } + + public reset(): void { + this.runtime = null; + this.currentModel = null; + } + + public async ensureRuntime(modelConfig: LocalSeparatorModelConfig, modelData: Uint8Array): Promise { + if (this.runtime && this.currentModel === modelConfig.filename) { + return this.runtime; + } + + if (!LocalOrtRuntimeManager.wasmPathsConfigured) { + ort.env.wasm.wasmPaths = { + mjs: ortWasmJsepMjsUrl, + wasm: ortWasmJsepUrl, + }; + localSeparatorLog('Configured ONNX Runtime wasm paths.', ort.env.wasm.wasmPaths); + LocalOrtRuntimeManager.wasmPathsConfigured = true; + } + + const providersToTry: Array<'webgpu' | 'wasm'> = []; + if (typeof navigator !== 'undefined' && 'gpu' in navigator) { + providersToTry.push('webgpu'); + localSeparatorLog('navigator.gpu is available, trying WebGPU first.'); + } else { + localSeparatorLog('navigator.gpu is not available. Falling back to CPU/wasm.'); + } + providersToTry.push('wasm'); + + let lastError: unknown = null; + for (const provider of providersToTry) { + try { + if (provider === 'webgpu' && ort.env?.webgpu) { + ort.env.webgpu.powerPreference = 'high-performance'; + localSeparatorLog('Using WebGPU power preference high-performance.'); + } + + const session = await ort.InferenceSession.create(modelData, { + executionProviders: [provider], + graphOptimizationLevel: 'all', + }); + + this.runtime = { provider, session }; + this.currentModel = modelConfig.filename; + this.onProviderChange(provider === 'wasm' ? 'cpu/wasm' : provider); + localSeparatorLog(`Using provider: ${provider}`); + return this.runtime; + } catch (error) { + lastError = error; + if (provider === 'webgpu') { + localSeparatorLog('WebGPU session creation failed. Falling back to CPU/wasm.', error); + this.onProviderChange('cpu/wasm fallback'); + } else { + localSeparatorLog(`Provider failed: ${provider}`, error); + } + } + } + + throw new Error(`Unable to create ONNX Runtime session. ${lastError ? String(lastError) : ''}`.trim()); + } +} diff --git a/src/util/localSeparatorShared.ts b/src/util/localSeparatorShared.ts new file mode 100644 index 0000000..b6549f1 --- /dev/null +++ b/src/util/localSeparatorShared.ts @@ -0,0 +1,275 @@ +import type { StereoChannels } from './localSeparatorTypes'; + +export function index4d(dims: number[], i0: number, i1: number, i2: number, i3: number): number { + return (((i0 * dims[1] + i1) * dims[2] + i2) * dims[3]) + i3; +} + +export function concatFloat32(parts: Float32Array[]): Float32Array { + const length = parts.reduce((sum, part) => sum + part.length, 0); + const output = new Float32Array(length); + let offset = 0; + for (const part of parts) { + output.set(part, offset); + offset += part.length; + } + return output; +} + +export function reflectPad(signal: Float32Array, leftPad: number, rightPad: number): Float32Array { + const result = new Float32Array(leftPad + signal.length + rightPad); + const last = signal.length - 1; + + for (let i = 0; i < leftPad; i += 1) { + result[i] = signal[leftPad - i]; + } + result.set(signal, leftPad); + for (let i = 0; i < rightPad; i += 1) { + result[leftPad + signal.length + i] = signal[last - 1 - i]; + } + + return result; +} + +export function nextPowerOfTwo(value: number): number { + let result = 1; + while (result < value) { + result <<= 1; + } + return result; +} + +export function scaleChannels(channels: StereoChannels, scale: number): StereoChannels { + return channels.map(channel => { + const output = new Float32Array(channel.length); + for (let i = 0; i < channel.length; i += 1) { + output[i] = channel[i] * scale; + } + return output; + }) as StereoChannels; +} + +export function negateArray(data: Float32Array): Float32Array { + const output = new Float32Array(data.length); + for (let i = 0; i < data.length; i += 1) { + output[i] = -data[i]; + } + return output; +} + +export function sliceChannels(channels: StereoChannels, startSample: number, endSample: number): StereoChannels { + return [ + channels[0].slice(startSample, endSample), + channels[1].slice(startSample, endSample), + ]; +} + +export function normalizeChannels( + channels: StereoChannels, + maxPeak: number, + minPeak: number | null, +): { channels: StereoChannels; originalPeak: number } { + let peak = 0; + for (const channel of channels) { + for (let i = 0; i < channel.length; i += 1) { + peak = Math.max(peak, Math.abs(channel[i])); + } + } + + if (peak === 0) { + return { channels, originalPeak: 0 }; + } + + let scale = 1; + if (peak > maxPeak) { + scale = maxPeak / peak; + } else if (minPeak !== null && peak < minPeak && minPeak > 0) { + scale = minPeak / peak; + } + + const normalized = channels.map(channel => { + const output = new Float32Array(channel.length); + for (let i = 0; i < channel.length; i += 1) { + output[i] = channel[i] * scale; + } + return output; + }) as StereoChannels; + + return { channels: normalized, originalPeak: peak }; +} + +export function createWindowCache(): { periodic: Map; symmetric: Map } { + return { + periodic: new Map(), + symmetric: new Map(), + }; +} + +export function getHannPeriodic( + length: number, + cache: { periodic: Map }, +): Float32Array { + const hit = cache.periodic.get(length); + if (hit) return hit; + + const window = new Float32Array(length); + for (let i = 0; i < length; i += 1) { + window[i] = 0.5 - 0.5 * Math.cos((2 * Math.PI * i) / length); + } + cache.periodic.set(length, window); + return window; +} + +export function getHanning( + length: number, + cache: { symmetric: Map }, +): Float32Array { + const hit = cache.symmetric.get(length); + if (hit) return hit; + + const window = new Float32Array(length); + if (length === 1) { + window[0] = 1; + } else { + for (let i = 0; i < length; i += 1) { + window[i] = 0.5 - 0.5 * Math.cos((2 * Math.PI * i) / (length - 1)); + } + } + cache.symmetric.set(length, window); + return window; +} + +export class FFT { + private readonly size: number; + + constructor(size: number) { + this.size = size; + } + + public transform(real: Float64Array, imag: Float64Array): void { + if (real.length !== imag.length || real.length !== this.size) { + throw new Error('FFT input shape mismatch.'); + } + + if ((this.size & (this.size - 1)) === 0) { + this.transformRadix2(real, imag); + } else { + this.transformBluestein(real, imag); + } + } + + public inverse(real: Float64Array, imag: Float64Array): void { + for (let i = 0; i < this.size; i += 1) { + imag[i] = -imag[i]; + } + this.transform(real, imag); + for (let i = 0; i < this.size; i += 1) { + real[i] /= this.size; + imag[i] = -imag[i] / this.size; + } + } + + private transformRadix2(real: Float64Array, imag: Float64Array): void { + const n = this.size; + const levels = Math.trunc(Math.log2(n)); + + for (let i = 0; i < n; i += 1) { + const j = reverseBits(i, levels); + if (j > i) { + [real[i], real[j]] = [real[j], real[i]]; + [imag[i], imag[j]] = [imag[j], imag[i]]; + } + } + + for (let size = 2; size <= n; size <<= 1) { + const halfsize = size >>> 1; + const tableStep = n / size; + for (let i = 0; i < n; i += size) { + for (let j = i, k = 0; j < i + halfsize; j += 1, k += tableStep) { + const angle = (2 * Math.PI * k) / n; + const cos = Math.cos(angle); + const sin = Math.sin(angle); + const tpre = real[j + halfsize] * cos + imag[j + halfsize] * sin; + const tpim = -real[j + halfsize] * sin + imag[j + halfsize] * cos; + real[j + halfsize] = real[j] - tpre; + imag[j + halfsize] = imag[j] - tpim; + real[j] += tpre; + imag[j] += tpim; + } + } + } + } + + private transformBluestein(real: Float64Array, imag: Float64Array): void { + const n = this.size; + const m = nextPowerOfTwo((n * 2) + 1); + const areal = new Float64Array(m); + const aimag = new Float64Array(m); + const breal = new Float64Array(m); + const bimag = new Float64Array(m); + const creal = new Float64Array(m); + const cimag = new Float64Array(m); + + for (let i = 0; i < n; i += 1) { + const angle = (Math.PI * ((i * i) % (n * 2))) / n; + const cos = Math.cos(angle); + const sin = Math.sin(angle); + areal[i] = real[i] * cos + imag[i] * sin; + aimag[i] = -real[i] * sin + imag[i] * cos; + breal[i] = cos; + bimag[i] = sin; + if (i !== 0) { + breal[m - i] = cos; + bimag[m - i] = sin; + } + } + + convolveComplex(areal, aimag, breal, bimag, creal, cimag); + + for (let i = 0; i < n; i += 1) { + const angle = (Math.PI * ((i * i) % (n * 2))) / n; + const cos = Math.cos(angle); + const sin = Math.sin(angle); + real[i] = (creal[i] * cos) + (cimag[i] * sin); + imag[i] = (-creal[i] * sin) + (cimag[i] * cos); + } + } +} + +function reverseBits(x: number, bits: number): number { + let y = 0; + for (let i = 0; i < bits; i += 1) { + y = (y << 1) | (x & 1); + x >>>= 1; + } + return y; +} + +function convolveComplex( + xreal: Float64Array, + ximag: Float64Array, + yreal: Float64Array, + yimag: Float64Array, + outreal: Float64Array, + outimag: Float64Array, +): void { + const n = xreal.length; + const fft = new FFT(n); + const xr = new Float64Array(xreal); + const xi = new Float64Array(ximag); + const yr = new Float64Array(yreal); + const yi = new Float64Array(yimag); + + fft.transform(xr, xi); + fft.transform(yr, yi); + + for (let i = 0; i < n; i += 1) { + const tempReal = (xr[i] * yr[i]) - (xi[i] * yi[i]); + const tempImag = (xi[i] * yr[i]) + (xr[i] * yi[i]); + xr[i] = tempReal; + xi[i] = tempImag; + } + + fft.inverse(xr, xi); + outreal.set(xr); + outimag.set(xi); +} diff --git a/src/util/localSeparatorTiming.ts b/src/util/localSeparatorTiming.ts new file mode 100644 index 0000000..f323a86 --- /dev/null +++ b/src/util/localSeparatorTiming.ts @@ -0,0 +1,52 @@ +export class LocalSeparatorTimingCollector { + private readonly label: string; + private readonly sections = new Map(); + private readonly startedAt = performance.now(); + + constructor(label = 'timing') { + this.label = label; + } + + private track(name: string, durationMs: number): void { + const section = this.sections.get(name) ?? { totalMs: 0, count: 0 }; + section.totalMs += durationMs; + section.count += 1; + this.sections.set(name, section); + } + + public async measureAsync(name: string, fn: () => Promise): Promise { + const started = performance.now(); + try { + return await fn(); + } finally { + this.track(name, performance.now() - started); + } + } + + public measureSync(name: string, fn: () => T): T { + const started = performance.now(); + try { + return fn(); + } finally { + this.track(name, performance.now() - started); + } + } + + public getSummary(extra: Record = {}): Record { + const sections: Record = {}; + for (const [name, section] of this.sections.entries()) { + sections[name] = { + totalMs: Number(section.totalMs.toFixed(2)), + count: section.count, + averageMs: Number((section.totalMs / Math.max(section.count, 1)).toFixed(2)), + }; + } + + return { + label: this.label, + totalMs: Number((performance.now() - this.startedAt).toFixed(2)), + sections, + ...extra, + }; + } +} diff --git a/src/util/localSeparatorTypes.ts b/src/util/localSeparatorTypes.ts new file mode 100644 index 0000000..a39326c --- /dev/null +++ b/src/util/localSeparatorTypes.ts @@ -0,0 +1,49 @@ +export type StereoChannels = [Float32Array, Float32Array]; + +export interface LocalSeparatorModelDefaults { + sampleRate: number; + hopLength: number; + segmentSize: number; + overlap: number; + batchSize: number; + enableDenoise: boolean; + invertUsingSpec: boolean; + normalizationThreshold: number; + amplificationThreshold: number; + matchMixOverlap: number; +} + +export interface LocalSeparatorModelMetadata { + compensate: number; + mdx_dim_f_set: number; + mdx_dim_t_set: number; + mdx_n_fft_scale_set: number; + primary_stem: string; +} + +export interface LocalSeparatorModelConfig { + filename: string; + displayName: string; + status: 'ready'; + defaults: LocalSeparatorModelDefaults; + metadata: LocalSeparatorModelMetadata; +} + +export interface LocalSeparatorProgress { + stage: string; + passLabel: string; + percent: number; + processedChunks: number; + totalChunks: number; +} + +export interface LocalRuntimeSupport { + webgpuExposed: boolean; +} + +export type LocalRuntimeProvider = 'webgpu' | 'wasm'; + +export interface LocalRuntimeState { + provider: LocalRuntimeProvider; + session: import('onnxruntime-web/webgpu').InferenceSession; +}