Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
621bd2784e | ||
|
|
4e56851264 | ||
|
|
8a84f6b48e | ||
|
|
2737d7b190 | ||
|
|
335af9fa4b | ||
|
|
accb0fd56c | ||
|
|
329a8826ad | ||
|
|
d8f885e61a | ||
|
|
82e0f78e21 | ||
|
|
17de6ca853 | ||
|
|
87aacaa903 | ||
|
|
4da4a7b498 | ||
|
|
f0cf78326c | ||
|
|
d387703101 | ||
|
|
1d4782516c | ||
|
|
ad8ad49d9c | ||
|
|
0b425c9bd1 | ||
|
|
a377c72fac | ||
|
|
ab3b9d1c3d | ||
|
|
c87923faeb | ||
|
|
64070adf77 | ||
|
|
b8ee1856c5 | ||
|
|
f973796f07 | ||
|
|
9a1297beaf | ||
|
|
93856c88eb | ||
|
|
75ed37ed3b | ||
|
|
08fb457a9d | ||
|
|
dfcce31318 | ||
|
|
499aac5e80 | ||
|
|
6760256cd7 |
@@ -2,10 +2,8 @@ name: Publish to Comfy registry
|
||||
on:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
tags:
|
||||
- '[0-9]+.[0-9]+.[0-9]+' # e.g. 1.2.3
|
||||
|
||||
jobs:
|
||||
publish-node:
|
||||
@@ -14,6 +12,7 @@ jobs:
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@main
|
||||
with:
|
||||
|
||||
@@ -14,3 +14,4 @@ __no__
|
||||
models/.catalog.csv
|
||||
models/*.yaml
|
||||
0LEEME
|
||||
sync.sh
|
||||
|
||||
@@ -41,3 +41,18 @@ repos:
|
||||
# "--check-hidden"
|
||||
]
|
||||
# You can create a .codespellignore file with one word per line for words to ignore.
|
||||
|
||||
# --- Version checking ---
|
||||
- repo: local
|
||||
hooks:
|
||||
- id: version-check
|
||||
name: check for version consistency
|
||||
# The command to execute. It's a Python script.
|
||||
entry: python3 tool/check_versions.py
|
||||
# Use 'system' to run it with the current environment's Python
|
||||
language: system
|
||||
# This hook doesn't need to run on every file.
|
||||
# It should run if either of the version files change.
|
||||
# This makes it very fast.
|
||||
files: ^(pyproject\.toml|src/nodes/__init__\.py)$
|
||||
# The regex `^...$` ensures it matches the full path from the repo root.
|
||||
|
||||
@@ -30,12 +30,19 @@ to keep the secondary vocals along with the instruments.
|
||||
|
||||
The objectives for these nodes are:
|
||||
|
||||
- Multiple stems (Vocals, Instruments, Drums, Bass, etc.)
|
||||
- Easy of use
|
||||
- Clear download (with progress and known destination)
|
||||
- Support for all possible input audio formats (mono/stereo, any sample rate, any batch size)
|
||||
- Good quality vs size, or you can choose better quality using Demucs
|
||||
- Reduced dependencies
|
||||
✅ Multiple stems (Vocals, Instruments, Drums, Bass, etc.)
|
||||
|
||||
✅ Easy of use
|
||||
|
||||
✅ Clear download (with progress and known destination)
|
||||
|
||||
✅ Support for all possible input audio formats (mono/stereo, any sample rate, any batch size)
|
||||
|
||||
✅ Good quality vs size, or you can choose better quality using Demucs
|
||||
|
||||
✅ Reduced dependencies, just my helpers when using ComfyUI
|
||||
|
||||
✅ Multiple examples
|
||||
|
||||
---
|
||||
|
||||
@@ -53,7 +60,9 @@ The objectives for these nodes are:
|
||||
* ✨ [Nodes](#-nodes)
|
||||
* [Vocals using MDX](#vocals-using-mdx)
|
||||
* [Demucs Audio Separator](#demucs-audio-separator)
|
||||
* 🖼️ [Examples](#️-examples)
|
||||
* 📝 [Usage Notes](#-usage-notes)
|
||||
* 📜 [Project History](#-project-history)
|
||||
* ⚖️ [License](#️-license)
|
||||
* 🙏 [Attributions](#-attributions)
|
||||
|
||||
@@ -71,10 +80,11 @@ or just do it manually:
|
||||
cd ComfyUI/custom_nodes/
|
||||
git clone https://github.com/set-soft/AudioSeparation
|
||||
```
|
||||
2. Restart ComfyUI.
|
||||
2. Install SeCoNoHe: `pip install seconohe`
|
||||
3. Restart ComfyUI.
|
||||
|
||||
The nodes should then appear under the "audio/separation" category in the "Add Node" menu.
|
||||
You don't need to install extra dependencies.
|
||||
SeCoNoHe are just a bunch of helpers I created with common functionality I use in my nodes.
|
||||
|
||||
### Command Line Tool
|
||||
|
||||
@@ -97,13 +107,7 @@ pip install -r requirements.txt
|
||||
4. Run the scripts like this:
|
||||
|
||||
```
|
||||
python3 tool/demix.py AUDIO_FILE
|
||||
````
|
||||
|
||||
or
|
||||
|
||||
```
|
||||
python tool/demix.py AUDIO_FILE
|
||||
tool/demix.py AUDIO_FILE
|
||||
````
|
||||
|
||||
You don't need to install it, you could even add a symlink in `/usr/bin`.
|
||||
@@ -250,6 +254,21 @@ And here is the Demucs node:
|
||||
- **Models:** Note that most models are a *bag of models*, this is four models working together.
|
||||
|
||||
|
||||
## 🖼️ Examples
|
||||
|
||||
Once installed the examples are available in the ComfyUI workflow templates, in the *audio-separation* section.
|
||||
|
||||
Note that we have two versions, the regular and the *quick* version. The *quick* version is ideal for quick tests, the input files
|
||||
are downloaded and, in most cases, only 10 seconds of audio are processed.
|
||||
|
||||
- [00_Vocals_simple.json](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/00_Vocals_simple.json) [Quick](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/00_Vocals_quick.json): Example to get vocals using MDX
|
||||
- [01_Vocals_Drums_Bass.json](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/01_Vocals_Drums_Bass.json) [Quick](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/01_Vocals_Drums_Bass_quick.json): Example to get vocals, drums, bass and others using MDX
|
||||
- [02_Batch.json](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/02_Batch.json) [Quick](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/02_Batch_quick.json): Shows how to apply MDX demix to a batch of audios, using *Audio Batch* nodes.
|
||||
- [03_Instrumental_keep.json](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/03_Instrumental_keep.json) [Quick](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/03_Instrumental_keep_quick.json): Shows how to extract vocals maintaining the same number of channels and sample rate, using *Audio Batch* nodes.
|
||||
- [04_Demucs.json](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/04_Demucs.json) [Quick](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/04_Demucs_quick.json): Separates vocals, drums, bass and others using Demucs, better quality.
|
||||
- [05_Demix_and_Remix.json](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/05_Demix_and_Remix.json) [Quick](https://raw.githubusercontent.com/set-soft/AudioSeparation/refs/heads/main/example_workflows/05_Demix_and_Remix_quick.json): Example to separate vocals and others to the left channel and drums and bass to the right channel, using *Audio Batch* nodes.
|
||||
|
||||
|
||||
## 📝 Usage Notes
|
||||
|
||||
- **AUDIO Type:** These nodes work with ComfyUI's standard "AUDIO" data type, which is a Python dictionary containing:
|
||||
@@ -263,6 +282,15 @@ And here is the Demucs node:
|
||||
- **No quantized Demucs:** These models just save download time, but pulls extra dependency (diffq), they are just lower quality versions of their non-quantized counterparts.
|
||||
|
||||
|
||||
## 📜 Project History
|
||||
|
||||
- 1.0.0 2025-07-02: Initial release. MDX-Net models support
|
||||
|
||||
- 1.1.0 2025-07-11: Demucs models support
|
||||
|
||||
- 1.1.1 2025-07-21: More examples. One more Demucs model
|
||||
|
||||
|
||||
## ⚖️ License
|
||||
|
||||
[GPL-3.0](LICENSE)
|
||||
|
||||
+6
-26
@@ -3,33 +3,13 @@
|
||||
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
# License: GPLv3
|
||||
# Project: ComfyUI-AudioSeparation
|
||||
import inspect
|
||||
import logging
|
||||
from .source.utils.misc import NODES_NAME
|
||||
from . import nodes # noqa: E402
|
||||
|
||||
init_logger = logging.getLogger(f"{NODES_NAME}.__init__")
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
# This is our first import so we initialize SeCoNoHe
|
||||
from .src.nodes import nodes, main_logger, __version__
|
||||
from seconohe.register_nodes import register_nodes
|
||||
from seconohe import JS_PATH
|
||||
|
||||
|
||||
def register_nodes(module):
|
||||
suffix = " " + module.SUFFIX if hasattr(module, "SUFFIX") else ""
|
||||
if suffix:
|
||||
suffix = " " + suffix
|
||||
for name, obj in inspect.getmembers(module):
|
||||
if not inspect.isclass(obj) or not hasattr(obj, "INPUT_TYPES"):
|
||||
continue
|
||||
assert hasattr(obj, "UNIQUE_NAME"), f"No name for {obj.__name__}"
|
||||
NODE_CLASS_MAPPINGS[obj.UNIQUE_NAME] = obj
|
||||
NODE_DISPLAY_NAME_MAPPINGS[obj.UNIQUE_NAME] = obj.DISPLAY_NAME + suffix
|
||||
|
||||
|
||||
register_nodes(nodes)
|
||||
|
||||
init_logger.info(f"Registering {len(NODE_CLASS_MAPPINGS)} node(s).")
|
||||
init_logger.debug(f"{list(NODE_DISPLAY_NAME_MAPPINGS.values())}")
|
||||
|
||||
WEB_DIRECTORY = "./js"
|
||||
NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS = register_nodes(main_logger, [nodes], version=__version__)
|
||||
WEB_DIRECTORY = JS_PATH
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 49 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 18 KiB |
@@ -0,0 +1,28 @@
|
||||
# Audio Separation assets
|
||||
|
||||
## Icon
|
||||
|
||||

|
||||
|
||||
prompt:
|
||||
|
||||
```
|
||||
Icon design, a single, vibrant white soundwave enters a sleek, crystalline triangular prism from the left.
|
||||
The prism refracts the soundwave, causing it to split into four distinct, brilliantly colored waveforms exiting to the right.
|
||||
One waveform is glowing red for vocals, one is electric blue for drums, one is deep green for bass, and one is bright yellow for other instruments.
|
||||
|
||||
Minimalist, vector art, graphic illustration, glowing neon lines, on a clean dark background. Visually compelling UI icon, 400x400.
|
||||
```
|
||||
|
||||
Model: HiDream I1
|
||||
|
||||
Version used:
|
||||
|
||||

|
||||
|
||||
|
||||
## Banner
|
||||
|
||||

|
||||
|
||||
Image composition
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 69 KiB |
File diff suppressed because one or more lines are too long
+1
@@ -0,0 +1 @@
|
||||
audioseparation_logo.jpg
|
||||
File diff suppressed because one or more lines are too long
@@ -0,0 +1 @@
|
||||
audioseparation_logo.jpg
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+1
@@ -0,0 +1 @@
|
||||
audioseparation_logo.jpg
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
@@ -0,0 +1 @@
|
||||
audioseparation_logo.jpg
|
||||
File diff suppressed because one or more lines are too long
+1
@@ -0,0 +1 @@
|
||||
audioseparation_logo.jpg
|
||||
File diff suppressed because one or more lines are too long
@@ -0,0 +1 @@
|
||||
audioseparation_logo.jpg
|
||||
File diff suppressed because one or more lines are too long
@@ -0,0 +1 @@
|
||||
audioseparation_logo.jpg
|
||||
File diff suppressed because one or more lines are too long
@@ -1,57 +0,0 @@
|
||||
// Copyright (c) 2025 Salvador E. Tropea
|
||||
// Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
// License: GPLv3
|
||||
// Project: ComfyUI-AudioSeparation
|
||||
|
||||
// This script adds an event named "set-audioseparation-node"
|
||||
// It can currently just modify a widget value for the current node
|
||||
|
||||
import { app } from "/scripts/app.js";
|
||||
|
||||
// Register a new extension
|
||||
app.registerExtension({
|
||||
name: "SET.AudioSeparation.NodeAdjust", // Unique name
|
||||
|
||||
// The setup function is executed when the extension is loaded
|
||||
setup() {
|
||||
// Add a listener for our custom event
|
||||
app.api.addEventListener("set-audioseparation-node", (event) => {
|
||||
// The data from Python is in event.detail
|
||||
const { action, arg1, arg2 } = event.detail;
|
||||
|
||||
// Find the node that is currently being executed
|
||||
const node = app.graph.getNodeById(app.runningNodeId);
|
||||
if (!node) {
|
||||
console.warn(`[SET.AudioSeparation] Could not find running node with ID: ${app.runningNodeId}`);
|
||||
return;
|
||||
}
|
||||
|
||||
// --- ACTION EXECUTED HERE ---
|
||||
switch (action) {
|
||||
case 'change_widget':
|
||||
// arg1 = widget name (e.g., "model")
|
||||
// arg2 = new value (e.g., "💾 My Awesome Model")
|
||||
|
||||
const widget = node.widgets.find(w => w.name === arg1);
|
||||
if (widget) {
|
||||
// This is the key part for combo boxes (dropdowns)
|
||||
// If the new value isn't in the list of options, add it first.
|
||||
if (!widget.options.values.includes(arg2)) {
|
||||
widget.options.values.push(arg2);
|
||||
}
|
||||
|
||||
// Set the widget value
|
||||
widget.setValue(arg2, node, app.canvas);
|
||||
} else {
|
||||
console.error(`[SET.AudioSeparation] Widget '${arg1}' not found on node ${node.id}`);
|
||||
}
|
||||
break;
|
||||
|
||||
// Other actions here in the future
|
||||
// case 'disable_widget':
|
||||
// ...
|
||||
// break;
|
||||
}
|
||||
});
|
||||
},
|
||||
});
|
||||
@@ -1,32 +0,0 @@
|
||||
// Copyright (c) 2025 Salvador E. Tropea
|
||||
// Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
// License: GPLv3
|
||||
// Project: ComfyUI-AudioSeparation
|
||||
|
||||
// This script adds an event named "set-audioseparation-toast"
|
||||
// Used to notify the user in the GUI using the Toast API
|
||||
|
||||
import { app } from "/scripts/app.js";
|
||||
|
||||
// Register a new extension
|
||||
app.registerExtension({
|
||||
name: "SET.AudioSeparation.ToastHandler", // Unique name
|
||||
|
||||
// The setup function is executed when the extension is loaded
|
||||
setup() {
|
||||
// Add a listener for our custom event
|
||||
app.api.addEventListener("set-audioseparation-toast", (event) => {
|
||||
// The data from Python is in event.detail
|
||||
const { message, summary, severity } = event.detail;
|
||||
|
||||
// Use the ComfyUI toast API to show the message
|
||||
// app.ui.toast.addMessage is the modern way to do this
|
||||
app.extensionManager.toast.add({
|
||||
severity: severity,
|
||||
summary: summary,
|
||||
detail: message,
|
||||
life: 6000
|
||||
});
|
||||
});
|
||||
},
|
||||
});
|
||||
@@ -543,6 +543,21 @@
|
||||
"primary_stem": "Vocals",
|
||||
"stages": 5
|
||||
},
|
||||
"5ee7b421b6228a90e63855b108ad830b": {
|
||||
"desc": "Hybrid Demucs MusDB High Train",
|
||||
"download": "Main/Demucs",
|
||||
"file_t": "safetensors",
|
||||
"model_t": "Demucs",
|
||||
"name": "hdemucs_high_trained.safetensors",
|
||||
"params": "83639368",
|
||||
"primary_stem": [
|
||||
"Vocals",
|
||||
"Drums",
|
||||
"Bass",
|
||||
"Other"
|
||||
],
|
||||
"use_demucs_pt_process": true
|
||||
},
|
||||
"5f6483271e1efb9bfb59e4a3e6d4d098": {
|
||||
"channels": 32,
|
||||
"compensate": 1.035,
|
||||
@@ -660,6 +675,21 @@
|
||||
"is_roformer": true,
|
||||
"model_type": "SCNet"
|
||||
},
|
||||
"6f82220b9fc8445331d3a246f4ab509a": {
|
||||
"desc": "Hybrid Demucs MusDB High Train",
|
||||
"download": "Torch/Demucs",
|
||||
"file_t": "pytorch",
|
||||
"model_t": "Demucs",
|
||||
"name": "hdemucs_high_trained.yaml",
|
||||
"params": 83639368,
|
||||
"primary_stem": [
|
||||
"Vocals",
|
||||
"Drums",
|
||||
"Bass",
|
||||
"Other"
|
||||
],
|
||||
"signatures": "[\"hdemucs_high_trained\"]"
|
||||
},
|
||||
"7205f4e53ec3226681ecd0c38e2fa747": {
|
||||
"desc": "Hybrid Transformer Demucs",
|
||||
"download": "Main/Demucs",
|
||||
|
||||
+17
-3
@@ -1,13 +1,27 @@
|
||||
[project]
|
||||
name = "audio-separation"
|
||||
description = "Audio separation (aka demixing) nodes, for Vocals, Instruments, Bass, Drums and Others. Using MDX-Net, no extra dependencies, support for batch and resample."
|
||||
version = "1.0.0"
|
||||
description = """
|
||||
Audio separation (aka demixing) nodes, for Vocals, Instruments, Bass, Drums and Others (experimental Piano and Guitar).
|
||||
Using MDX-Net and Demucs, no extra dependencies, support for batch and resample.
|
||||
Choose between High Quality and Speed. All safetensor models (No ONNX, No PyTorch)
|
||||
"""
|
||||
# Inconsistent mechanism needed by comfy-cli, no dynamic variables
|
||||
version = "1.1.3"
|
||||
# Deprecated mechanism, comfy-cli doesn't support SPDX
|
||||
license = { file = "LICENSE" }
|
||||
dependencies = []
|
||||
# Not really used, ComfyUI-Manager doesn't use it
|
||||
# dependencies = ["seconohe>=1.0.2"]
|
||||
# So we do it in the reverse way ...
|
||||
dynamic = ["dependencies"]
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/set-soft/AudioSeparation"
|
||||
|
||||
[tool.setuptools.dynamic]
|
||||
dependencies = {file = ["requirements.txt"]}
|
||||
|
||||
[tool.comfy]
|
||||
PublisherId = "set-soft"
|
||||
DisplayName = "Audio Separation (Demix)"
|
||||
Icon = "https://raw.githubusercontent.com/set-soft/AudioSeparation/main/assets/AudioSeparation_400.jpg"
|
||||
Banner = "https://raw.githubusercontent.com/set-soft/AudioSeparation/main/assets/audioseparation_logo_21_9.jpg"
|
||||
|
||||
@@ -3,6 +3,7 @@ torchaudio
|
||||
numpy
|
||||
safetensors
|
||||
tqdm
|
||||
seconohe>=1.0.2
|
||||
# Optionals:
|
||||
# requests
|
||||
# colorama
|
||||
|
||||
@@ -1,220 +0,0 @@
|
||||
# File under the MIT license, see https://github.com/adefossez/julius/LICENSE for details.
|
||||
# Author: adefossez, 2020
|
||||
"""
|
||||
Differentiable, Pytorch based resampling.
|
||||
Implementation of Julius O. Smith algorithm for resampling.
|
||||
See https://ccrma.stanford.edu/~jos/resample/ for details.
|
||||
This implementation is specially optimized for when new_sr / old_sr is a fraction
|
||||
with a small numerator and denominator when removing the gcd (e.g. new_sr = 700, old_sr = 500).
|
||||
|
||||
Very similar to [bmcfee/resampy](https://github.com/bmcfee/resampy) except this implementation
|
||||
is optimized for the case mentioned before, while resampy is slower but more general.
|
||||
|
||||
"""
|
||||
import math
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
from torch.nn import functional as F
|
||||
|
||||
|
||||
def sinc(x: torch.Tensor):
|
||||
"""
|
||||
Implementation of sinc, i.e. sin(x) / x
|
||||
|
||||
__Warning__: the input is not multiplied by `pi`!
|
||||
"""
|
||||
return torch.where(x == 0, torch.tensor(1., device=x.device, dtype=x.dtype), torch.sin(x) / x)
|
||||
|
||||
|
||||
class ResampleFrac(torch.nn.Module):
|
||||
"""
|
||||
Resampling from the sample rate `old_sr` to `new_sr`.
|
||||
"""
|
||||
def __init__(self, old_sr: int, new_sr: int, zeros: int = 24, rolloff: float = 0.945):
|
||||
"""
|
||||
Args:
|
||||
old_sr (int): sample rate of the input signal x.
|
||||
new_sr (int): sample rate of the output.
|
||||
zeros (int): number of zero crossing to keep in the sinc filter.
|
||||
rolloff (float): use a lowpass filter that is `rolloff * new_sr / 2`,
|
||||
to ensure sufficient margin due to the imperfection of the FIR filter used.
|
||||
Lowering this value will reduce anti-aliasing, but will reduce some of the
|
||||
highest frequencies.
|
||||
|
||||
Shape:
|
||||
|
||||
- Input: `[*, T]`
|
||||
- Output: `[*, T']` with `T' = int(new_sr * T / old_sr)
|
||||
|
||||
|
||||
.. caution::
|
||||
After dividing `old_sr` and `new_sr` by their GCD, both should be small
|
||||
for this implementation to be fast.
|
||||
|
||||
>>> import torch
|
||||
>>> resample = ResampleFrac(4, 5)
|
||||
>>> x = torch.randn(1000)
|
||||
>>> print(len(resample(x)))
|
||||
1250
|
||||
"""
|
||||
super().__init__()
|
||||
if not isinstance(old_sr, int) or not isinstance(new_sr, int):
|
||||
raise ValueError("old_sr and new_sr should be integers")
|
||||
gcd = math.gcd(old_sr, new_sr)
|
||||
self.old_sr = old_sr // gcd
|
||||
self.new_sr = new_sr // gcd
|
||||
self.zeros = zeros
|
||||
self.rolloff = rolloff
|
||||
|
||||
self._init_kernels()
|
||||
|
||||
def _init_kernels(self):
|
||||
if self.old_sr == self.new_sr:
|
||||
return
|
||||
|
||||
kernels = []
|
||||
sr = min(self.new_sr, self.old_sr)
|
||||
# rolloff will perform antialiasing filtering by removing the highest frequencies.
|
||||
# At first I thought I only needed this when downsampling, but when upsampling
|
||||
# you will get edge artifacts without this, the edge is equivalent to zero padding,
|
||||
# which will add high freq artifacts.
|
||||
sr *= self.rolloff
|
||||
|
||||
# The key idea of the algorithm is that x(t) can be exactly reconstructed from x[i] (tensor)
|
||||
# using the sinc interpolation formula:
|
||||
# x(t) = sum_i x[i] sinc(pi * old_sr * (i / old_sr - t))
|
||||
# We can then sample the function x(t) with a different sample rate:
|
||||
# y[j] = x(j / new_sr)
|
||||
# or,
|
||||
# y[j] = sum_i x[i] sinc(pi * old_sr * (i / old_sr - j / new_sr))
|
||||
|
||||
# We see here that y[j] is the convolution of x[i] with a specific filter, for which
|
||||
# we take an FIR approximation, stopping when we see at least `zeros` zeros crossing.
|
||||
# But y[j+1] is going to have a different set of weights and so on, until y[j + new_sr].
|
||||
# Indeed:
|
||||
# y[j + new_sr] = sum_i x[i] sinc(pi * old_sr * ((i / old_sr - (j + new_sr) / new_sr))
|
||||
# = sum_i x[i] sinc(pi * old_sr * ((i - old_sr) / old_sr - j / new_sr))
|
||||
# = sum_i x[i + old_sr] sinc(pi * old_sr * (i / old_sr - j / new_sr))
|
||||
# so y[j+new_sr] uses the same filter as y[j], but on a shifted version of x by `old_sr`.
|
||||
# This will explain the F.conv1d after, with a stride of old_sr.
|
||||
self._width = math.ceil(self.zeros * self.old_sr / sr)
|
||||
# If old_sr is still big after GCD reduction, most filters will be very unbalanced, i.e.,
|
||||
# they will have a lot of almost zero values to the left or to the right...
|
||||
# There is probably a way to evaluate those filters more efficiently, but this is kept for
|
||||
# future work.
|
||||
idx = torch.arange(-self._width, self._width + self.old_sr).float()
|
||||
for i in range(self.new_sr):
|
||||
t = (-i/self.new_sr + idx/self.old_sr) * sr
|
||||
t = t.clamp_(-self.zeros, self.zeros)
|
||||
t *= math.pi
|
||||
window = torch.cos(t/self.zeros/2)**2
|
||||
kernel = sinc(t) * window
|
||||
# Renormalize kernel to ensure a constant signal is preserved.
|
||||
kernel.div_(kernel.sum())
|
||||
kernels.append(kernel)
|
||||
|
||||
self.register_buffer("kernel", torch.stack(kernels).view(self.new_sr, 1, -1))
|
||||
|
||||
def forward(self, x: torch.Tensor, output_length: Optional[int] = None, full: bool = False):
|
||||
"""
|
||||
Resample x.
|
||||
Args:
|
||||
x (Tensor): signal to resample, time should be the last dimension
|
||||
output_length (None or int): This can be set to the desired output length
|
||||
(last dimension). Allowed values are between 0 and
|
||||
ceil(length * new_sr / old_sr). When None (default) is specified, the
|
||||
floored output length will be used. In order to select the largest possible
|
||||
size, use the `full` argument.
|
||||
full (bool): return the longest possible output from the input. This can be useful
|
||||
if you chain resampling operations, and want to give the `output_length` only
|
||||
for the last one, while passing `full=True` to all the other ones.
|
||||
"""
|
||||
if self.old_sr == self.new_sr:
|
||||
return x
|
||||
shape = x.shape
|
||||
length = x.shape[-1]
|
||||
x = x.reshape(-1, length)
|
||||
x = F.pad(x[:, None], (self._width, self._width + self.old_sr), mode='replicate')
|
||||
ys = F.conv1d(x, self.kernel, stride=self.old_sr) # type: ignore
|
||||
y = ys.transpose(1, 2).reshape(list(shape[:-1]) + [-1])
|
||||
|
||||
float_output_length = torch.as_tensor(self.new_sr * length / self.old_sr)
|
||||
max_output_length = torch.ceil(float_output_length).long()
|
||||
default_output_length = torch.floor(float_output_length).long()
|
||||
|
||||
if output_length is None:
|
||||
applied_output_length = max_output_length if full else default_output_length
|
||||
elif output_length < 0 or output_length > max_output_length:
|
||||
raise ValueError(f"output_length must be between 0 and {max_output_length}")
|
||||
else:
|
||||
applied_output_length = torch.tensor(output_length)
|
||||
if full:
|
||||
raise ValueError("You cannot pass both full=True and output_length")
|
||||
return y[..., :applied_output_length] # type: ignore
|
||||
|
||||
|
||||
def resample_frac(x: torch.Tensor, old_sr: int, new_sr: int,
|
||||
zeros: int = 24, rolloff: float = 0.945,
|
||||
output_length: Optional[int] = None, full: bool = False):
|
||||
"""
|
||||
Functional version of `ResampleFrac`, refer to its documentation for more information.
|
||||
|
||||
..warning::
|
||||
If you call repeatidly this functions with the same sample rates, then the
|
||||
resampling kernel will be recomputed every time. For best performance, you should use
|
||||
and cache an instance of `ResampleFrac`.
|
||||
"""
|
||||
return ResampleFrac(old_sr, new_sr, zeros, rolloff).to(x)(x, output_length, full)
|
||||
|
||||
|
||||
# Easier implementations for downsampling and upsampling by a factor of 2
|
||||
# Kept for testing and reference
|
||||
|
||||
def _kernel_upsample2_downsample2(zeros):
|
||||
# Kernel for upsampling and downsampling by a factor of 2. Interestingly,
|
||||
# it is the same kernel used for both.
|
||||
win = torch.hann_window(4 * zeros + 1, periodic=False)
|
||||
winodd = win[1::2]
|
||||
t = torch.linspace(-zeros + 0.5, zeros - 0.5, 2 * zeros)
|
||||
t *= math.pi
|
||||
kernel = (sinc(t) * winodd).view(1, 1, -1)
|
||||
return kernel
|
||||
|
||||
|
||||
def _upsample2(x, zeros=24):
|
||||
"""
|
||||
Upsample x by a factor of two. The output will be exactly twice as long as the input.
|
||||
Args:
|
||||
x (Tensor): signal to upsample, time should be the last dimension
|
||||
zeros (int): number of zero crossing to keep in the sinc filter.
|
||||
|
||||
This function is kept only for reference, you should use the more generic `resample_frac`
|
||||
one. This function does not perform anti-aliasing filtering.
|
||||
"""
|
||||
*other, time = x.shape
|
||||
kernel = _kernel_upsample2_downsample2(zeros).to(x)
|
||||
out = F.conv1d(x.view(-1, 1, time), kernel, padding=zeros)[..., 1:].view(*other, time)
|
||||
y = torch.stack([x, out], dim=-1)
|
||||
return y.view(*other, -1)
|
||||
|
||||
|
||||
def _downsample2(x, zeros=24):
|
||||
"""
|
||||
Downsample x by a factor of two. The output length is half of the input, ceiled.
|
||||
Args:
|
||||
x (Tensor): signal to downsample, time should be the last dimension
|
||||
zeros (int): number of zero crossing to keep in the sinc filter.
|
||||
|
||||
This function is kept only for reference, you should use the more generic `resample_frac`
|
||||
one. This function does not perform anti-aliasing filtering.
|
||||
"""
|
||||
if x.shape[-1] % 2 != 0:
|
||||
x = F.pad(x, (0, 1))
|
||||
xeven = x[..., ::2]
|
||||
xodd = x[..., 1::2]
|
||||
*other, time = xodd.shape
|
||||
kernel = _kernel_upsample2_downsample2(zeros).to(x)
|
||||
out = xeven + F.conv1d(xodd.view(-1, 1, time), kernel, padding=zeros)[..., :-1].view(
|
||||
*other, time)
|
||||
return out.view(*other, -1).mul(0.5)
|
||||
@@ -1,113 +0,0 @@
|
||||
# Copyright Jonathan Hartley 2013. BSD 3-Clause license, see LICENSE file.
|
||||
'''
|
||||
This module generates ANSI character codes to printing colors to terminals.
|
||||
See: http://en.wikipedia.org/wiki/ANSI_escape_code
|
||||
'''
|
||||
import sys
|
||||
import os
|
||||
|
||||
CSI = '\033['
|
||||
OSC = '\033]'
|
||||
BEL = '\a'
|
||||
is_a_tty = sys.stderr.isatty() and os.name == 'posix'
|
||||
|
||||
|
||||
def code_to_chars(code):
|
||||
return CSI + str(code) + 'm' if is_a_tty else ''
|
||||
|
||||
|
||||
def set_title(title):
|
||||
return OSC + '2;' + title + BEL
|
||||
|
||||
|
||||
def clear_screen(mode=2):
|
||||
return CSI + str(mode) + 'J'
|
||||
|
||||
|
||||
def clear_line(mode=2):
|
||||
return CSI + str(mode) + 'K'
|
||||
|
||||
|
||||
class AnsiCodes(object):
|
||||
def __init__(self):
|
||||
# the subclasses declare class attributes which are numbers.
|
||||
# Upon instantiation we define instance attributes, which are the same
|
||||
# as the class attributes but wrapped with the ANSI escape sequence
|
||||
for name in dir(self):
|
||||
if not name.startswith('_'):
|
||||
value = getattr(self, name)
|
||||
setattr(self, name, code_to_chars(value))
|
||||
|
||||
|
||||
class AnsiCursor(object):
|
||||
def UP(self, n=1):
|
||||
return CSI + str(n) + 'A'
|
||||
|
||||
def DOWN(self, n=1):
|
||||
return CSI + str(n) + 'B'
|
||||
|
||||
def FORWARD(self, n=1):
|
||||
return CSI + str(n) + 'C'
|
||||
|
||||
def BACK(self, n=1):
|
||||
return CSI + str(n) + 'D'
|
||||
|
||||
def POS(self, x=1, y=1):
|
||||
return CSI + str(y) + ';' + str(x) + 'H'
|
||||
|
||||
|
||||
class AnsiFore(AnsiCodes):
|
||||
BLACK = 30
|
||||
RED = 31
|
||||
GREEN = 32
|
||||
YELLOW = 33
|
||||
BLUE = 34
|
||||
MAGENTA = 35
|
||||
CYAN = 36
|
||||
WHITE = 37
|
||||
RESET = 39
|
||||
|
||||
# These are fairly well supported, but not part of the standard.
|
||||
LIGHTBLACK_EX = 90
|
||||
LIGHTRED_EX = 91
|
||||
LIGHTGREEN_EX = 92
|
||||
LIGHTYELLOW_EX = 93
|
||||
LIGHTBLUE_EX = 94
|
||||
LIGHTMAGENTA_EX = 95
|
||||
LIGHTCYAN_EX = 96
|
||||
LIGHTWHITE_EX = 97
|
||||
|
||||
|
||||
class AnsiBack(AnsiCodes):
|
||||
BLACK = 40
|
||||
RED = 41
|
||||
GREEN = 42
|
||||
YELLOW = 43
|
||||
BLUE = 44
|
||||
MAGENTA = 45
|
||||
CYAN = 46
|
||||
WHITE = 47
|
||||
RESET = 49
|
||||
|
||||
# These are fairly well supported, but not part of the standard.
|
||||
LIGHTBLACK_EX = 100
|
||||
LIGHTRED_EX = 101
|
||||
LIGHTGREEN_EX = 102
|
||||
LIGHTYELLOW_EX = 103
|
||||
LIGHTBLUE_EX = 104
|
||||
LIGHTMAGENTA_EX = 105
|
||||
LIGHTCYAN_EX = 106
|
||||
LIGHTWHITE_EX = 107
|
||||
|
||||
|
||||
class AnsiStyle(AnsiCodes):
|
||||
BRIGHT = 1
|
||||
DIM = 2
|
||||
NORMAL = 22
|
||||
RESET_ALL = 0
|
||||
|
||||
|
||||
Fore = AnsiFore()
|
||||
Back = AnsiBack()
|
||||
Style = AnsiStyle()
|
||||
Cursor = AnsiCursor()
|
||||
@@ -1,44 +0,0 @@
|
||||
# Copyright (c) 2025 Salvador E. Tropea
|
||||
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
# License: GPLv3
|
||||
# Project: ComfyUI-AudioSeparation
|
||||
#
|
||||
# ComfyUI Node actions
|
||||
import logging
|
||||
# ComfyUI imports
|
||||
try:
|
||||
from server import PromptServer
|
||||
with_comfy = True
|
||||
except Exception:
|
||||
with_comfy = False
|
||||
# Local imports
|
||||
from .misc import NODES_NAME
|
||||
|
||||
logger = logging.getLogger(f"{NODES_NAME}.comfy_node_action")
|
||||
|
||||
|
||||
def send_node_action(action: str, arg1: str = None, arg2: str = None, sid: str = None):
|
||||
"""
|
||||
Sends a node action event to the ComfyUI client.
|
||||
|
||||
Args:
|
||||
action (str): Action to be performed.
|
||||
arg1 (str): First argument
|
||||
arg2 (str): Second argument
|
||||
sid (str, optional): The session ID of the client to send to.
|
||||
If None, broadcasts to all clients. Defaults to None.
|
||||
"""
|
||||
if not with_comfy:
|
||||
return
|
||||
try:
|
||||
PromptServer.instance.send_sync(
|
||||
"set-audioseparation-node", # This is our custom event name
|
||||
{
|
||||
'action': action,
|
||||
'arg1': arg1,
|
||||
'arg2': arg2
|
||||
},
|
||||
sid
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"when trying to use ComfyUI PromptServer: {e}")
|
||||
@@ -1,46 +0,0 @@
|
||||
# Copyright (c) 2025 Salvador E. Tropea
|
||||
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
# License: GPLv3
|
||||
# Project: ComfyUI-AudioSeparation
|
||||
#
|
||||
# ComfyUI Toast API messages
|
||||
# Original code from Gemini 2.5 Pro, which was really outdated
|
||||
# Took ideas from Easy Use nodes and looking at ComfyUI code
|
||||
import logging
|
||||
# ComfyUI imports
|
||||
try:
|
||||
from server import PromptServer
|
||||
with_comfy = True
|
||||
except Exception:
|
||||
with_comfy = False
|
||||
# Local imports
|
||||
from .misc import NODES_NAME
|
||||
|
||||
logger = logging.getLogger(f"{NODES_NAME}.comfy_notification")
|
||||
|
||||
|
||||
def send_toast_notification(message: str, summary: str = "Warning", severity: str = "warn", sid: str = None):
|
||||
"""
|
||||
Sends a toast notification event to the ComfyUI client.
|
||||
|
||||
Args:
|
||||
message (str): The message content of the toast.
|
||||
severity (str): The type of toast. Can be 'success' | 'info' | 'warn' | 'error' | 'secondary' | 'contrast'
|
||||
summary (str): Short explanation
|
||||
sid (str, optional): The session ID of the client to send to.
|
||||
If None, broadcasts to all clients. Defaults to None.
|
||||
"""
|
||||
if not with_comfy:
|
||||
return
|
||||
try:
|
||||
PromptServer.instance.send_sync(
|
||||
"set-audioseparation-toast", # This is our custom event name
|
||||
{
|
||||
'message': message,
|
||||
'summary': summary,
|
||||
'severity': severity
|
||||
},
|
||||
sid
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"when trying to use ComfyUI PromptServer: {e}")
|
||||
@@ -1,206 +0,0 @@
|
||||
# Copyright (c) 2025 Salvador E. Tropea
|
||||
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
# License: GPLv3
|
||||
# Project: ComfyUI-AudioSeparation
|
||||
#
|
||||
# Model downloader w/TQDM and ComfyUI progress
|
||||
# Original code from Gemini 2.5 Pro
|
||||
import logging
|
||||
import os
|
||||
# Requests is better than the core Python urllib, and is a really common package
|
||||
# But we don't really need it. Lets make it optional:
|
||||
try:
|
||||
import requests
|
||||
with_requests = True
|
||||
except Exception:
|
||||
with_requests = False
|
||||
import urllib
|
||||
from tqdm import tqdm
|
||||
# ComfyUI imports
|
||||
try:
|
||||
import comfy.utils
|
||||
with_comfy = True
|
||||
except Exception:
|
||||
with_comfy = False
|
||||
# Local imports
|
||||
from .misc import NODES_NAME
|
||||
|
||||
logger = logging.getLogger(f"{NODES_NAME}.downloader")
|
||||
|
||||
|
||||
def download_model_requests(url: str, save_dir: str, file_name: str):
|
||||
"""
|
||||
Downloads a file from a URL with progress bars for both console and ComfyUI.
|
||||
|
||||
Args:
|
||||
url (str): The direct download URL for the file.
|
||||
save_dir (str): The directory where the file will be saved.
|
||||
file_name (str): The name of the file to be saved on disk.
|
||||
"""
|
||||
full_path = os.path.join(save_dir, file_name)
|
||||
|
||||
# Ensure the save directory exists
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
try:
|
||||
# Use a streaming request to handle large files and get content length
|
||||
with requests.get(url, stream=True, timeout=10) as r:
|
||||
r.raise_for_status() # Raise an exception for bad status codes (4xx or 5xx)
|
||||
|
||||
# Get total file size from headers
|
||||
total_size_in_bytes = int(r.headers.get('content-length', 0))
|
||||
block_size = 1024 # 1 Kibibyte
|
||||
|
||||
# --- Setup Progress Bars ---
|
||||
# Console progress bar using tqdm
|
||||
progress_bar_console = tqdm(
|
||||
total=total_size_in_bytes,
|
||||
unit='iB',
|
||||
unit_scale=True,
|
||||
desc=f"Downloading {file_name}"
|
||||
)
|
||||
|
||||
# ComfyUI progress bar
|
||||
progress_bar_ui = comfy.utils.ProgressBar(total_size_in_bytes) if with_comfy else None
|
||||
|
||||
# --- Download Loop ---
|
||||
downloaded_size = 0
|
||||
with open(full_path, 'wb') as f:
|
||||
for chunk in r.iter_content(chunk_size=block_size):
|
||||
if chunk: # filter out keep-alive new chunks
|
||||
chunk_size = len(chunk)
|
||||
|
||||
# Update console progress bar
|
||||
progress_bar_console.update(chunk_size)
|
||||
|
||||
# Update ComfyUI progress bar
|
||||
downloaded_size += chunk_size
|
||||
if progress_bar_ui:
|
||||
progress_bar_ui.update(chunk_size) # ProgressBar takes absolute value, but update is incremental
|
||||
|
||||
# Write chunk to file
|
||||
f.write(chunk)
|
||||
|
||||
# --- Cleanup ---
|
||||
progress_bar_console.close()
|
||||
|
||||
# Final check to see if download was complete
|
||||
if total_size_in_bytes != 0 and progress_bar_console.n != total_size_in_bytes:
|
||||
logger.error("Download failed: Size mismatch.")
|
||||
# Optional: remove partial file
|
||||
# os.remove(full_path)
|
||||
raise IOError(f"Download failed for {file_name}. Expected {total_size_in_bytes} but got "
|
||||
f"{progress_bar_console.n}")
|
||||
|
||||
return full_path
|
||||
|
||||
except requests.exceptions.RequestException as e:
|
||||
logger.error(f"Network error while downloading {file_name}: {e}")
|
||||
# Clean up partial file if it exists
|
||||
if os.path.exists(full_path):
|
||||
try:
|
||||
os.remove(full_path)
|
||||
except OSError:
|
||||
pass
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"An error occurred during download: {e}")
|
||||
if os.path.exists(full_path):
|
||||
try:
|
||||
os.remove(full_path)
|
||||
except OSError:
|
||||
pass
|
||||
raise
|
||||
|
||||
|
||||
# A simple version implemented using the Python urllib
|
||||
class Downloader:
|
||||
def __init__(self, model_path, model_name):
|
||||
self.model_path = model_path
|
||||
self.model_name = model_name
|
||||
self.model_full_name = os.path.join(self.model_path, self.model_name)
|
||||
# Ensure the directory for the model_path exists before __init__ if used elsewhere
|
||||
# or create it at the start of download_model
|
||||
|
||||
# A TQDM helper class for urlretrieve reporthook
|
||||
# This is a common pattern for this use case.
|
||||
class TqdmUpTo(tqdm):
|
||||
"""
|
||||
Provides `update_to(block_num, block_size, total_size)`
|
||||
and updates the TQDM bar.
|
||||
"""
|
||||
def __init__(self, unit, unit_scale, unit_divisor, miniters, desc):
|
||||
super().__init__(unit=unit, unit_scale=unit_scale, unit_divisor=unit_divisor, miniters=miniters, desc=desc)
|
||||
self.ui_bar = None
|
||||
self.total = None
|
||||
|
||||
def update_to(self, block_num=1, block_size=1, total_size=None):
|
||||
"""
|
||||
block_num : int, optional
|
||||
Number of blocks transferred so far [default: 1].
|
||||
block_size : int, optional
|
||||
Size of each block (in tqdm units) [default: 1].
|
||||
total_size : int, optional
|
||||
Total size (in tqdm units). If [default: None] remains unchanged.
|
||||
"""
|
||||
if total_size is not None and self.total is None:
|
||||
self.total = total_size
|
||||
# ComfyUI progress bar
|
||||
if self.ui_bar is None and with_comfy:
|
||||
self.ui_bar = comfy.utils.ProgressBar(total_size)
|
||||
# self.update() will take the *difference* from the last call.
|
||||
# So we pass the number of new blocks * block_size.
|
||||
# Since block_num is cumulative, we calculate the new amount.
|
||||
chunk_size = block_num * block_size - self.n
|
||||
self.update(chunk_size) # self.n is current progress
|
||||
if self.ui_bar:
|
||||
self.ui_bar.update(chunk_size) # ProgressBar takes absolute value, but update is incremental
|
||||
|
||||
def download_model(self, url: str):
|
||||
try:
|
||||
# Ensure the directory exists
|
||||
# Use or '.' for current dir if dirname is empty
|
||||
os.makedirs(self.model_path or '.', exist_ok=True)
|
||||
|
||||
# Get filename for tqdm description
|
||||
filename = self.model_name
|
||||
|
||||
# Use TqdmUpTo as a context manager
|
||||
with self.TqdmUpTo(unit='iB', unit_scale=True, unit_divisor=1024, miniters=1,
|
||||
desc=f"Downloading {filename}") as t:
|
||||
# urlretrieve(url, filename=None, reporthook=None, data=None)
|
||||
# reporthook is called with (block_num, block_size, total_size)
|
||||
urllib.request.urlretrieve(url, self.model_full_name, reporthook=t.update_to)
|
||||
# The 'with' statement ensures t.close() is called.
|
||||
|
||||
return filename
|
||||
|
||||
except urllib.error.URLError as e: # More specific exception for network issues
|
||||
# Clean up partially downloaded file if an error occurs
|
||||
if os.path.exists(self.model_full_name):
|
||||
os.remove(self.model_full_name)
|
||||
raise Exception(f"An error occurred while downloading the model (URL Error): {e.reason} from {url}")
|
||||
|
||||
except Exception as e:
|
||||
# Clean up partially downloaded file if an error occurs
|
||||
if os.path.exists(self.model_full_name):
|
||||
os.remove(self.model_full_name)
|
||||
raise Exception(f"An unexpected error occurred while downloading the model: {e}")
|
||||
|
||||
|
||||
def download_model_urllib(url: str, save_dir: str, file_name: str):
|
||||
return Downloader(save_dir, file_name).download_model(url)
|
||||
|
||||
|
||||
def download_model(url: str, save_dir: str, file_name: str, force_urllib: bool = False):
|
||||
logger.info(f"Downloading model: {file_name}")
|
||||
logger.info(f"Source URL: {url}")
|
||||
full_name = os.path.join(save_dir, file_name)
|
||||
logger.info(f"Destination: {full_name}")
|
||||
|
||||
if with_requests and not force_urllib:
|
||||
download_model_requests(url, save_dir, file_name)
|
||||
else:
|
||||
download_model_urllib(url, save_dir, file_name)
|
||||
|
||||
logger.info(f"Successfully downloaded {full_name}")
|
||||
return full_name
|
||||
@@ -1,104 +0,0 @@
|
||||
# Copyright (c) 2025 Salvador E. Tropea
|
||||
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
# License: GPLv3
|
||||
# Project: ComfyUI-AudioSeparation
|
||||
import os
|
||||
import sys
|
||||
import logging
|
||||
from .misc import NODES_NAME, NODES_DEBUG_VAR
|
||||
|
||||
no_colorama = False
|
||||
try:
|
||||
from colorama import init as colorama_init, Fore, Back, Style
|
||||
except ImportError:
|
||||
no_colorama = True
|
||||
# If colorama isn't installed use an ANSI basic replacement
|
||||
if no_colorama:
|
||||
from .ansi import Fore, Back, Style # noqa: F811
|
||||
else:
|
||||
colorama_init()
|
||||
|
||||
# Used for tools
|
||||
standalone_mode = False
|
||||
|
||||
white = Fore.WHITE + Style.BRIGHT
|
||||
yellow = Fore.YELLOW + Style.BRIGHT
|
||||
red = Fore.RED + Style.BRIGHT
|
||||
red_alarm = Fore.RED + Back.WHITE + Style.BRIGHT
|
||||
cyan = Fore.CYAN + Style.BRIGHT
|
||||
reset = Style.RESET_ALL
|
||||
# format = "%(asctime)s - %(name)s - %(levelname)s - %(message)s "
|
||||
# "(%(filename)s:%(lineno)d)"
|
||||
format = f"[{NODES_NAME} %(levelname)s] %(message)s (%(name)s - %(filename)s:%(lineno)d)"
|
||||
format_simple = f"[{NODES_NAME}] %(message)s"
|
||||
FORMATS = {
|
||||
logging.DEBUG: cyan + format + reset,
|
||||
logging.INFO: white + format_simple + reset,
|
||||
logging.WARNING: yellow + format + reset,
|
||||
logging.ERROR: red + format + reset,
|
||||
logging.CRITICAL: red_alarm + format + reset
|
||||
}
|
||||
format = "[%(levelname)s] %(message)s (%(name)s - %(filename)s:%(lineno)d)"
|
||||
format_simple = "%(message)s"
|
||||
if not sys.stdout.isatty():
|
||||
white = yellow = red = red_alarm = cyan = reset = ""
|
||||
FORMATS_STANDALONE = {
|
||||
logging.DEBUG: cyan + format + reset,
|
||||
logging.INFO: white + format_simple + reset,
|
||||
logging.WARNING: yellow + format + reset,
|
||||
logging.ERROR: red + format + reset,
|
||||
logging.CRITICAL: red_alarm + format + reset
|
||||
}
|
||||
|
||||
|
||||
class CustomFormatter(logging.Formatter):
|
||||
"""Logging Formatter to add colors"""
|
||||
|
||||
def __init__(self):
|
||||
super(logging.Formatter, self).__init__()
|
||||
|
||||
def format(self, record):
|
||||
formats = FORMATS_STANDALONE if standalone_mode else FORMATS
|
||||
log_fmt = formats.get(record.levelno)
|
||||
formatter = logging.Formatter(log_fmt)
|
||||
return formatter.format(record)
|
||||
|
||||
|
||||
# Create a new logger
|
||||
logger = logging.getLogger(NODES_NAME)
|
||||
logger.propagate = False
|
||||
|
||||
# Add handler if we don't have one.
|
||||
if not logger.handlers:
|
||||
handler = logging.StreamHandler(sys.stdout)
|
||||
handler.setFormatter(CustomFormatter())
|
||||
logger.addHandler(handler)
|
||||
|
||||
# ######################
|
||||
# Logger setup
|
||||
# ######################
|
||||
# 1. Determine the ComfyUI global log level (influenced by --verbose)
|
||||
main_logger = logger
|
||||
comfy_root_logger = logging.getLogger('comfy')
|
||||
effective_comfy_level = logging.getLogger().getEffectiveLevel()
|
||||
# 2. Check our custom environment variable for more verbosity
|
||||
try:
|
||||
nodes_debug_env = int(os.environ.get(NODES_DEBUG_VAR, "0"))
|
||||
except ValueError:
|
||||
nodes_debug_env = 0
|
||||
# 3. Set node's logger level
|
||||
if nodes_debug_env:
|
||||
main_logger.setLevel(logging.DEBUG - (nodes_debug_env - 1))
|
||||
final_level_str = f"DEBUG (due to {NODES_DEBUG_VAR}={nodes_debug_env})"
|
||||
else:
|
||||
main_logger.setLevel(effective_comfy_level)
|
||||
final_level_str = logging.getLevelName(effective_comfy_level) + " (matching ComfyUI global)"
|
||||
_initial_setup_logger = logging.getLogger(NODES_NAME + ".setup") # A temporary logger for this message
|
||||
_initial_setup_logger.debug(f"{NODES_NAME} logger level set to: {final_level_str}")
|
||||
|
||||
|
||||
def logger_set_standalone(args):
|
||||
verbose = args.verbose
|
||||
main_logger.setLevel(logging.DEBUG - (verbose - 1) if verbose else logging.INFO)
|
||||
global standalone_mode
|
||||
standalone_mode = True
|
||||
@@ -1,128 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
# Copyright (c) 2025 Salvador E. Tropea
|
||||
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
# License: CC BY-NC-SA 4.0
|
||||
# Project: ComfyUI-Float_Optimized
|
||||
import contextlib # For context manager
|
||||
import logging
|
||||
import torch
|
||||
try:
|
||||
import comfy.model_management as mm
|
||||
with_comfy = True
|
||||
except Exception:
|
||||
with_comfy = False
|
||||
from .misc import NODES_NAME
|
||||
|
||||
logger = logging.getLogger(f"{NODES_NAME}.torch")
|
||||
|
||||
|
||||
def get_torch_device_options():
|
||||
# We always have CPU
|
||||
default = "cpu"
|
||||
options = [default]
|
||||
# Do we have CUDA?
|
||||
if torch.cuda.is_available():
|
||||
default = "cuda"
|
||||
options.append(default)
|
||||
for i in range(torch.cuda.device_count()):
|
||||
options.append(f"cuda:{i}") # Specific CUDA devices
|
||||
# Is this a Mac?
|
||||
if torch.backends.mps.is_available() and torch.backends.mps.is_built():
|
||||
options.append("mps")
|
||||
if default == "cpu":
|
||||
default = "mps"
|
||||
return options, default
|
||||
|
||||
|
||||
def get_offload_device():
|
||||
return mm.unet_offload_device() if with_comfy else torch.device("cpu")
|
||||
|
||||
|
||||
def get_canonical_device(device: str | torch.device) -> torch.device:
|
||||
"""Converts a device string or object into a canonical torch.device object with an explicit index."""
|
||||
if not isinstance(device, torch.device):
|
||||
device = torch.device(device)
|
||||
|
||||
# If it's a CUDA device and no index is specified, get the default one.
|
||||
if device.type == 'cuda' and device.index is None:
|
||||
# NOTE: This adds a dependency on torch.cuda.current_device()
|
||||
# The first solution is often better as it doesn't need this.
|
||||
return torch.device(f'cuda:{torch.cuda.current_device()}')
|
||||
return device
|
||||
|
||||
|
||||
# ##################################################################################
|
||||
# # Helper for inference (Target device, offload, eval, no_grad and cuDNN Benchmark)
|
||||
# ##################################################################################
|
||||
|
||||
@contextlib.contextmanager
|
||||
def model_to_target(model):
|
||||
"""
|
||||
Consolidated context manager for model device placement and inference state.
|
||||
|
||||
- Moves the model to its designated `model.target_device`.
|
||||
- Sets `torch.backends.cudnn.benchmark` based on `model.cudnn_benchmark_setting` if available.
|
||||
- Sets the model to `eval()` mode.
|
||||
- Wraps the operation in a `torch.no_grad()` context.
|
||||
- Offloads the model to the CPU (`mm.unet_offload_device()`) afterwards.
|
||||
"""
|
||||
if not isinstance(model, torch.nn.Module):
|
||||
with torch.no_grad():
|
||||
yield # The code inside the 'with' statement runs here
|
||||
return
|
||||
|
||||
# 1. Determine target device from the model object
|
||||
try:
|
||||
target_device = model.target_device
|
||||
assert isinstance(target_device, torch.device)
|
||||
except Exception as e:
|
||||
logger.warning(f"model_to_target: Could not get 'target_device' from model ({e}). "
|
||||
"Defaulting to model's current device.")
|
||||
target_device = next(model.parameters()).device
|
||||
|
||||
# 2. Get CUDNN benchmark setting from the model object (optional)
|
||||
# Use hasattr as this is an optional setting that not all models might have.
|
||||
cudnn_benchmark_enabled = None # Default is to keep the current setting
|
||||
if hasattr(model, 'cudnn_benchmark_setting'):
|
||||
cudnn_benchmark_enabled = model.cudnn_benchmark_setting
|
||||
|
||||
original_device = next(model.parameters()).device
|
||||
original_cudnn_benchmark_state = None
|
||||
is_cuda_target = target_device.type == 'cuda'
|
||||
|
||||
try:
|
||||
# 3. Manage cuDNN benchmark state
|
||||
if (cudnn_benchmark_enabled is not None and is_cuda_target and hasattr(torch.backends, 'cudnn') and
|
||||
torch.backends.cudnn.is_available()):
|
||||
if torch.backends.cudnn.benchmark != cudnn_benchmark_enabled:
|
||||
original_cudnn_benchmark_state = torch.backends.cudnn.benchmark
|
||||
torch.backends.cudnn.benchmark = cudnn_benchmark_enabled
|
||||
logger.debug(f"Temporarily set cuDNN benchmark to {torch.backends.cudnn.benchmark}")
|
||||
|
||||
# 4. Move model to target device if not already there
|
||||
if original_device != target_device:
|
||||
logger.debug(f"Moving model from `{original_device}` to target device `{target_device}` for inference.")
|
||||
model.to(target_device)
|
||||
|
||||
# 5. Set to eval mode and disable gradients for the operation
|
||||
model.eval()
|
||||
with torch.no_grad():
|
||||
yield # The code inside the 'with' statement runs here
|
||||
|
||||
finally:
|
||||
# 6. Restore original cuDNN benchmark state
|
||||
if original_cudnn_benchmark_state is not None:
|
||||
# This check is sufficient because it will only be not None if we set it inside the try block
|
||||
torch.backends.cudnn.benchmark = original_cudnn_benchmark_state
|
||||
logger.debug(f"Restored cuDNN benchmark to {original_cudnn_benchmark_state}")
|
||||
|
||||
# 7. Offload model back to CPU
|
||||
if with_comfy:
|
||||
offload_device = get_offload_device()
|
||||
current_device_after_yield = next(model.parameters()).device
|
||||
if current_device_after_yield != offload_device:
|
||||
logger.debug(f"Offloading model from `{current_device_after_yield}` to offload device `{offload_device}`.")
|
||||
model.to(offload_device)
|
||||
# Clear cache if we were on a CUDA device
|
||||
if 'cuda' in str(current_device_after_yield):
|
||||
torch.cuda.empty_cache()
|
||||
@@ -0,0 +1,13 @@
|
||||
# Copyright (c) 2025 Salvador E. Tropea
|
||||
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
# License: GPL-3.0
|
||||
# Project: ComfyUI-AudioSeparation
|
||||
from seconohe.logger import initialize_logger
|
||||
|
||||
|
||||
__version__ = "1.1.3"
|
||||
__copyright__ = "Copyright © 2025 Salvador E. Tropea / Instituto Nacional de Tecnología Industrial"
|
||||
__license__ = "License GPLv3+: GNU GPL version 3 or later <https://gnu.org/licenses/gpl.html>"
|
||||
__author__ = "Salvador E. Tropea"
|
||||
NODES_NAME = "AudioSeparation"
|
||||
main_logger = initialize_logger(NODES_NAME)
|
||||
@@ -10,10 +10,11 @@
|
||||
import os
|
||||
import csv
|
||||
import logging
|
||||
from seconohe.logger import debugl
|
||||
import sys
|
||||
|
||||
from .hash import get_hash
|
||||
from ..utils.misc import NODES_NAME, debugl
|
||||
from .. import NODES_NAME
|
||||
|
||||
# Set up the logger as specified
|
||||
logger = logging.getLogger(f"{NODES_NAME}.hash_dir")
|
||||
@@ -132,7 +133,7 @@ def hash_dir(directory_path: str) -> dict:
|
||||
|
||||
# ==============================================================================
|
||||
# Command-Line Tool for Testing and Operation
|
||||
# python -m source.db.hash_dir
|
||||
# python -m src.nodes.db.hash_dir
|
||||
# ==============================================================================
|
||||
if __name__ == "__main__":
|
||||
# Local imports to avoid top-level pollution
|
||||
@@ -9,7 +9,7 @@ import logging
|
||||
# Local imports
|
||||
from .models_db import download_model
|
||||
from ..inference.get_model import get_model
|
||||
from ..utils.misc import NODES_NAME
|
||||
from .. import NODES_NAME
|
||||
|
||||
logger = logging.getLogger(f"{NODES_NAME}.load_model")
|
||||
|
||||
@@ -9,9 +9,8 @@ import json
|
||||
import logging
|
||||
import os
|
||||
from pathlib import Path
|
||||
from ..utils.misc import NODES_NAME
|
||||
from ..utils.downloader import download_model as download_model_basic
|
||||
from ..utils.comfy_notification import send_toast_notification
|
||||
from seconohe.downloader import download_file as download_model_basic
|
||||
from .. import NODES_NAME
|
||||
from .hash_dir import hash_dir
|
||||
from .hash import is_hash, get_hash
|
||||
|
||||
@@ -33,7 +32,7 @@ def get_db_filename(provided=None):
|
||||
script_dir = Path(__file__).resolve().parent
|
||||
|
||||
# Build the path to the JSON file: go up one level, then into models/
|
||||
json_path = script_dir / ".." / ".." / "models" / "uvr_model_data.json"
|
||||
json_path = script_dir / ".." / ".." / ".." / "models" / "uvr_model_data.json"
|
||||
try:
|
||||
return json_path.resolve().relative_to(Path.cwd())
|
||||
except ValueError:
|
||||
@@ -244,7 +243,7 @@ def get_download_url(data):
|
||||
except KeyError:
|
||||
return None
|
||||
try:
|
||||
return os.path.join(KNOWN_SOURCES[dn_t], name)
|
||||
return KNOWN_SOURCES[dn_t] + '/' + name
|
||||
except KeyError:
|
||||
logger.error(f"Unknown download source `{dn_t}`")
|
||||
return None
|
||||
@@ -273,15 +272,13 @@ def download_model(data, models_dir):
|
||||
|
||||
# Download the file
|
||||
name = data['name']
|
||||
send_toast_notification(f"Downloading `{name}`", "Download")
|
||||
try:
|
||||
fname = download_model_basic(url, models_dir, name)
|
||||
fname = download_model_basic(logger, url, models_dir, name)
|
||||
# Mark it as downloaded
|
||||
data['model_path'] = fname
|
||||
if data['indicator']:
|
||||
data['indicator'] = ICON_DOWNLOADED
|
||||
# Notify the user
|
||||
send_toast_notification("Finished downloading", "Download", 'success')
|
||||
return fname
|
||||
except Exception as e:
|
||||
raise ValueError(f"Failed to download {name} from {url}\n{e}")
|
||||
@@ -11,8 +11,8 @@ import typing as tp
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.nn import functional as F
|
||||
import torchaudio
|
||||
|
||||
from .resample_frac import resample_frac # From julius
|
||||
from .demucs_code import capture_init, center_trim, unfold
|
||||
from .CrossTransformerEncoder import LayerScale
|
||||
|
||||
@@ -373,6 +373,8 @@ class Demucs(nn.Module):
|
||||
if rescale:
|
||||
rescale_module(self, reference=rescale)
|
||||
|
||||
self.upsampler = self.downsampler = None
|
||||
|
||||
def valid_length(self, length):
|
||||
"""
|
||||
Return the nearest valid length to use with the model so that
|
||||
@@ -413,7 +415,16 @@ class Demucs(nn.Module):
|
||||
x = F.pad(x, (delta // 2, delta - delta // 2))
|
||||
|
||||
if self.resample:
|
||||
x = resample_frac(x, 1, 2)
|
||||
if self.upsampler is None:
|
||||
# Create the resamplers as instance attributes.
|
||||
# They will be automatically moved to the correct device when model.to(device) is called.
|
||||
self.upsampler = torchaudio.transforms.Resample(orig_freq=self.samplerate,
|
||||
new_freq=2 * self.samplerate,
|
||||
lowpass_filter_width=24).to(x.device)
|
||||
self.downsampler = torchaudio.transforms.Resample(orig_freq=2 * self.samplerate,
|
||||
new_freq=self.samplerate,
|
||||
lowpass_filter_width=24).to(x.device)
|
||||
x = self.upsampler(x)
|
||||
|
||||
saved = []
|
||||
for encode in self.encoder:
|
||||
@@ -429,7 +440,7 @@ class Demucs(nn.Module):
|
||||
x = decode(x + skip)
|
||||
|
||||
if self.resample:
|
||||
x = resample_frac(x, 2, 1)
|
||||
x = self.downsampler(x)
|
||||
x = x * std + mean
|
||||
x = center_trim(x, length)
|
||||
x = x.view(x.size(0), len(self.sources), self.audio_channels, x.size(-1))
|
||||
@@ -70,7 +70,7 @@ class ScaledEmbedding(nn.Module):
|
||||
class HEncLayer(nn.Module):
|
||||
def __init__(self, chin, chout, kernel_size=8, stride=4, norm_groups=1, empty=False,
|
||||
freq=True, dconv=True, norm=True, context=0, dconv_kw={}, pad=True,
|
||||
rewrite=True):
|
||||
rewrite=True, force_norm_in_last=False):
|
||||
"""Encoder layer. This used both by the time and the frequency branch.
|
||||
|
||||
Args:
|
||||
@@ -109,6 +109,9 @@ class HEncLayer(nn.Module):
|
||||
pad = [pad, 0]
|
||||
klass = nn.Conv2d
|
||||
self.conv = klass(chin, chout, kernel_size, stride, pad)
|
||||
if force_norm_in_last:
|
||||
# PyTorch Audio uses it on last (empty) layer
|
||||
self.norm1 = norm_fn(chout)
|
||||
if self.empty:
|
||||
return
|
||||
self.norm1 = norm_fn(chout)
|
||||
@@ -257,7 +260,7 @@ class MultiWrap(nn.Module):
|
||||
class HDecLayer(nn.Module):
|
||||
def __init__(self, chin, chout, last=False, kernel_size=8, stride=4, norm_groups=1, empty=False,
|
||||
freq=True, dconv=True, norm=True, context=1, dconv_kw={}, pad=True,
|
||||
context_freq=True, rewrite=True):
|
||||
context_freq=True, rewrite=True, force_norm_in_last=False):
|
||||
"""
|
||||
Same as HEncLayer but for decoder. See `HEncLayer` for documentation.
|
||||
"""
|
||||
@@ -397,6 +400,7 @@ class HDemucs(nn.Module):
|
||||
# Normalization
|
||||
norm_starts=4,
|
||||
norm_groups=4,
|
||||
force_norm_in_last=False,
|
||||
# DConv residual branch
|
||||
dconv_mode=1,
|
||||
dconv_depth=2,
|
||||
@@ -533,6 +537,7 @@ class HDemucs(nn.Module):
|
||||
kwt['kernel_size'] = kernel_size
|
||||
kwt['stride'] = stride
|
||||
kwt['pad'] = True
|
||||
kwt['force_norm_in_last'] = force_norm_in_last
|
||||
kw_dec = dict(kw)
|
||||
multi = False
|
||||
if multi_freqs and index < multi_freqs_depth:
|
||||
@@ -6,7 +6,9 @@
|
||||
# Wrappers for the model and inference
|
||||
import logging
|
||||
import math
|
||||
from seconohe.torch import model_to_target, get_offload_device
|
||||
import torch
|
||||
import tqdm
|
||||
# ComfyUI imports
|
||||
try:
|
||||
import comfy.utils
|
||||
@@ -16,10 +18,11 @@ except Exception:
|
||||
# Local imports
|
||||
from .stft import stft_chunk_process, stft_get_chunks
|
||||
from ..db.load_model import load_model
|
||||
from ..utils.misc import NODES_NAME
|
||||
from ..utils.torch import model_to_target, get_offload_device
|
||||
from .. import NODES_NAME
|
||||
from .demucs_api import apply_model, BagOfModels
|
||||
|
||||
from torchaudio.transforms import Fade
|
||||
|
||||
logger = logging.getLogger(f"{NODES_NAME}.demixer")
|
||||
SAMPLE_RATE = 44100
|
||||
|
||||
@@ -135,6 +138,62 @@ def get_steps_for_demucs(model, wav, segment, shifts, overlap):
|
||||
return math.ceil(wav.shape[-1] / stride) * (shifts + 1)
|
||||
|
||||
|
||||
def separate_sources(
|
||||
model: torch.nn.Module,
|
||||
mix: torch.Tensor,
|
||||
sample_rate: int,
|
||||
segment: float = 10.0,
|
||||
overlap: float = 0.1,
|
||||
device: torch.device = None,
|
||||
chunk_fade_shape: str = "linear",
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
From: https://pytorch.org/audio/stable/tutorials/hybrid_demucs_tutorial.html
|
||||
|
||||
Apply model to a given mixture. Use fade, and add segments together in order to add model segment by segment.
|
||||
|
||||
Args:
|
||||
segment (int): segment length in seconds
|
||||
device (torch.device, str, or None): if provided, device on which to
|
||||
execute the computation, otherwise `mix.device` is assumed.
|
||||
When `device` is different from `mix.device`, only local computations will
|
||||
be on `device`, while the entire tracks will be stored on `mix.device`.
|
||||
"""
|
||||
batch, channels, length = mix.shape
|
||||
|
||||
chunk_len = int(sample_rate * segment * (1 + overlap))
|
||||
start = 0
|
||||
end = chunk_len
|
||||
overlap_frames = overlap * sample_rate
|
||||
fade = Fade(fade_in_len=0, fade_out_len=int(overlap_frames), fade_shape=chunk_fade_shape)
|
||||
chunks = math.ceil((length - overlap_frames) / chunk_len)
|
||||
# Progress bars
|
||||
progress_bar_console = tqdm.tqdm(total=chunks)
|
||||
if with_comfy:
|
||||
comfy_progress_bar = comfy.utils.ProgressBar(chunks)
|
||||
|
||||
final = torch.zeros(batch, len(model.sources), channels, length, device=device)
|
||||
|
||||
while start < length - overlap_frames:
|
||||
chunk = mix[:, :, start:end]
|
||||
progress_bar_console.update(1)
|
||||
if with_comfy:
|
||||
comfy_progress_bar.update(1)
|
||||
with torch.no_grad():
|
||||
out = model.forward(chunk)
|
||||
out = fade(out)
|
||||
final[:, :, :, start:end] += out
|
||||
if start == 0:
|
||||
fade.fade_in_len = int(overlap_frames)
|
||||
start += int(chunk_len - overlap_frames)
|
||||
else:
|
||||
start += chunk_len
|
||||
end += chunk_len
|
||||
if end >= length:
|
||||
fade.fade_out_len = 0
|
||||
return final
|
||||
|
||||
|
||||
class DemixerDemucs(DemixerGeneric):
|
||||
def __init__(self, d, device, models_dir):
|
||||
super().__init__(d, device, models_dir)
|
||||
@@ -194,21 +253,38 @@ class DemixerDemucs(DemixerGeneric):
|
||||
forced_segment = None
|
||||
else:
|
||||
logger.debug(f"Using user provided segment size {forced_segment} s")
|
||||
if with_comfy:
|
||||
comfy_progress_bar = comfy.utils.ProgressBar(self.get_steps(waveform_tensor, forced_segment, shifts, overlap))
|
||||
|
||||
with model_to_target(model):
|
||||
separated_tensors = apply_model(
|
||||
model,
|
||||
input_tensor_on_device,
|
||||
device=self.device,
|
||||
segment=forced_segment,
|
||||
shifts=shifts + 1, # Shifts 0 is disabled, 1 is just one pass, 2 is 2 passes
|
||||
overlap=overlap,
|
||||
split=True, # Enable chunking
|
||||
progress=True, # Show a progress bar in the console
|
||||
callback=lambda x: (x.get('state') == 'end') and comfy_progress_bar.update(1) if with_comfy else None,
|
||||
)
|
||||
if self.d.get('use_demucs_pt_process', False):
|
||||
# This is for the model from torchaudio, using the Demucs code we get some strange noises
|
||||
# Using their example they aren't produced
|
||||
model.target_device = self.device
|
||||
logger.debug("Using PyTorch Audio chunking for old model")
|
||||
with model_to_target(logger, model):
|
||||
separated_tensors = separate_sources(
|
||||
model,
|
||||
input_tensor_on_device,
|
||||
model.samplerate,
|
||||
segment=forced_segment or 16.0,
|
||||
overlap=overlap,
|
||||
device=self.device,
|
||||
chunk_fade_shape="half_sine"
|
||||
)
|
||||
else:
|
||||
if with_comfy:
|
||||
comfy_progress_bar = comfy.utils.ProgressBar(self.get_steps(waveform_tensor, forced_segment, shifts,
|
||||
overlap))
|
||||
with model_to_target(logger, model):
|
||||
separated_tensors = apply_model(
|
||||
model,
|
||||
input_tensor_on_device,
|
||||
device=self.device,
|
||||
segment=forced_segment,
|
||||
shifts=shifts + 1, # Shifts 0 is disabled, 1 is just one pass, 2 is 2 passes
|
||||
overlap=overlap,
|
||||
split=True, # Enable chunking
|
||||
progress=True, # Show a progress bar in the console
|
||||
callback=lambda x: (x.get('state') == 'end') and comfy_progress_bar.update(1) if with_comfy else None,
|
||||
)
|
||||
|
||||
# Move the final result tensor back to the CPU before creating the output dicts.
|
||||
# This is good practice to free up VRAM for subsequent nodes.
|
||||
@@ -265,19 +341,23 @@ class DemixerDemucs(DemixerGeneric):
|
||||
output_dict = {
|
||||
"waveform": output_tensor,
|
||||
"sample_rate": model.samplerate,
|
||||
"stem": stem_name.capitalize()
|
||||
"stem": stem_name.capitalize(),
|
||||
"generated": True,
|
||||
}
|
||||
final_outputs.append(output_dict)
|
||||
else:
|
||||
logger.debug(f"Stem '{stem_name}' not produced by model. Returning silence.")
|
||||
# Create a silent tensor with the correct shape, device, and dtype
|
||||
silent_waveform = torch.zeros(ref_batch_size, ref_channels, ref_samples, dtype=torch.float32)
|
||||
if not input_was_batched:
|
||||
silent_waveform = silent_waveform.squeeze(0)
|
||||
|
||||
# Wrap the silent tensor in the ComfyUI AUDIO dictionary format
|
||||
silent_dict = {
|
||||
"waveform": silent_waveform,
|
||||
"sample_rate": model.samplerate,
|
||||
"stem": stem_name.capitalize()
|
||||
"stem": stem_name.capitalize(),
|
||||
"generated": False,
|
||||
}
|
||||
final_outputs.append(silent_dict)
|
||||
|
||||
@@ -25,9 +25,9 @@ from .HTDemucs import HTDemucs
|
||||
from .demucs_code import center_trim, DummyPoolExecutor
|
||||
# AudioSeparation stuff
|
||||
import logging
|
||||
from ..utils.misc import NODES_NAME
|
||||
from .. import NODES_NAME
|
||||
|
||||
logger = logging.getLogger(f"{NODES_NAME}.demixer")
|
||||
logger = logging.getLogger(f"{NODES_NAME}.demucs_api")
|
||||
Model = tp.Union[Demucs, HDemucs, HTDemucs]
|
||||
|
||||
|
||||
@@ -0,0 +1,322 @@
|
||||
from fractions import Fraction
|
||||
import logging
|
||||
from .. import NODES_NAME
|
||||
|
||||
rlogger = logging.getLogger(f"{NODES_NAME}.demucs_log")
|
||||
|
||||
|
||||
class DemucsModelInfo(object):
|
||||
def __init__(self, index, klass_name: str, kwargs: dict, logger, weights, extra=False, sig=None):
|
||||
super().__init__()
|
||||
self.index = index
|
||||
self.kwargs = kwargs
|
||||
self.logger = logger
|
||||
self.extra = extra
|
||||
self.extra_indent = ""
|
||||
self.weights = weights
|
||||
self.signature = sig
|
||||
|
||||
# Always log these fundamental parameters
|
||||
sr = kwargs.get('samplerate', 44100)
|
||||
if sr != 44100:
|
||||
rlogger.warning("Model not configured for 44.1 kHz sample rate")
|
||||
a_ch = kwargs.get('audio_channels', 2)
|
||||
if a_ch != 2:
|
||||
rlogger.warning("Model not configured for stereo")
|
||||
|
||||
if klass_name == "HTDemucs":
|
||||
self.htdemucs()
|
||||
elif klass_name == "HDemucs":
|
||||
self.hdemucs()
|
||||
elif klass_name == "Demucs":
|
||||
self.demucs()
|
||||
else:
|
||||
logger.warning(f"No specific logger for model class: {klass_name}. "
|
||||
"Displaying raw kwargs.")
|
||||
for key, value in kwargs.items():
|
||||
logger(f" - {key}: {value}")
|
||||
|
||||
def get(self, key, default=None):
|
||||
return self.kwargs.get(key, default)
|
||||
|
||||
def log_type(self, name):
|
||||
start = "" if self.index < 0 else f"{self.index+1}. "
|
||||
msg = f" {start}Type: {name}"
|
||||
if self.signature:
|
||||
msg += f" [{self.signature}]"
|
||||
self.logger(msg)
|
||||
|
||||
def _log_param(self, param_name, default, description="", unit="", indent=" ", can_skip=False):
|
||||
"""
|
||||
Logs a parameter if its value is different from the default, or if it's a key parameter.
|
||||
|
||||
Args:
|
||||
param_name (str): The name of the parameter to check.
|
||||
default: The default value for this parameter.
|
||||
description (str): A user-friendly description of the parameter.
|
||||
unit (str): An optional unit to display after the value (e.g., 'Hz').
|
||||
indent (str): The indentation string for the log message.
|
||||
"""
|
||||
value = self.get(param_name, default)
|
||||
|
||||
# We log if the value is not the default, or if it's a fundamental parameter.
|
||||
is_default = (value == default)
|
||||
is_important = param_name in ['sources', 'segment']
|
||||
|
||||
if is_default and not is_important and not self.extra:
|
||||
return None
|
||||
if can_skip and is_default:
|
||||
return None
|
||||
|
||||
if param_name == 'sources' and self.weights:
|
||||
value = [s if w == 1.0 else ('' if not w else f'{w}*{s}') for s, w in zip(value, self.weights)]
|
||||
|
||||
if unit == '%':
|
||||
value *= 100
|
||||
unit_str = f" {unit}" if unit else ""
|
||||
desc_str = description if description else param_name.capitalize().replace('_', ' ')
|
||||
indent += self.extra_indent
|
||||
value_str = f"{value.numerator}/{value.denominator}" if isinstance(value, Fraction) else str(value)
|
||||
n = (39 - len(desc_str) - len(value_str) - len(unit_str))*" "
|
||||
return f" {indent}- {desc_str}: {value_str}{unit_str} {n}({param_name})"
|
||||
|
||||
def log_param(self, param_name, default, description="", unit="", indent=" ", can_skip=False):
|
||||
res = self._log_param(param_name, default, description, unit, indent, can_skip)
|
||||
if res is not None:
|
||||
self.logger(res)
|
||||
|
||||
def add(self, param_name, default, description="", unit="", indent=" ", can_skip=False):
|
||||
res = self._log_param(param_name, default, description, unit, indent, can_skip)
|
||||
if res is not None:
|
||||
self.params.append(res)
|
||||
|
||||
def reset(self):
|
||||
self.params = []
|
||||
|
||||
def sub_section(self, name):
|
||||
self.logger(f" {name}:")
|
||||
|
||||
def section(self, name):
|
||||
self.logger(" " + "-" * 40)
|
||||
self.sub_section(name)
|
||||
|
||||
def flush(self, name, is_sub=False):
|
||||
if self.params:
|
||||
if is_sub:
|
||||
self.sub_section(self.extra_indent + name)
|
||||
else:
|
||||
self.section(name)
|
||||
for p in self.params:
|
||||
self.logger(p)
|
||||
|
||||
def structure(self, ch=64, depth=6, with_lstm=False, with_ch_tm=False):
|
||||
self.reset()
|
||||
self.add('channels', ch, "Initial hidden channels")
|
||||
self.add('depth', depth, "Number of U-Net layers")
|
||||
self.add('growth', 2.0, "Channel growth factor per layer")
|
||||
self.add('rewrite', True, "Use 1x1 convolutions in blocks")
|
||||
if with_lstm:
|
||||
self.add('lstm_layers', 0, "Number of main LSTM layers", can_skip=True)
|
||||
if with_ch_tm:
|
||||
self.add('channels_time', None, "Specific channels for time branch", can_skip=True)
|
||||
self.flush("Structure")
|
||||
|
||||
def convolutions(self, advanced=False):
|
||||
self.reset()
|
||||
self.add('kernel_size', 8)
|
||||
self.add('stride', 4)
|
||||
if advanced:
|
||||
self.add('time_stride', 2, "Final time layer stride")
|
||||
self.add('context', 1, "Decoder context window size")
|
||||
if advanced:
|
||||
self.add('context_enc', 0, "Encoder context window size")
|
||||
self.flush("Convolutions")
|
||||
|
||||
def normalization(self):
|
||||
self.reset()
|
||||
self.add('norm_starts', 4, "Start at layer")
|
||||
self.add('norm_groups', 4, "Number of groups")
|
||||
self.flush("Normalization")
|
||||
|
||||
def dconv(self, full=True):
|
||||
if self.get('dconv_mode', 1) <= 0:
|
||||
return
|
||||
self.reset()
|
||||
where = ['', 'In encoder', 'In decoder', 'In encoder and decoder'][self.get('dconv_mode', 1)]
|
||||
self.add('dconv_mode', 1, where)
|
||||
self.add('dconv_depth', 2, "Number of layers in DConv branch")
|
||||
if full:
|
||||
comp = 4
|
||||
init = 1e-4
|
||||
else:
|
||||
comp = 8
|
||||
init = 1e-3
|
||||
self.add('dconv_comp', comp, "Channel compression factor")
|
||||
self.add('dconv_init', init, "Initial scale")
|
||||
if full:
|
||||
self.add('dconv_attn', 4, "Layer to start attention in DConv")
|
||||
self.add('dconv_lstm', 4, "Layer to start LSTM in DConv")
|
||||
self.flush("DConv Residual Branch")
|
||||
|
||||
def stft(self):
|
||||
self.reset()
|
||||
self.add('nfft', 4096, "Frequency Bins")
|
||||
# Decode the method
|
||||
cac = self.get('cac')
|
||||
niters = self.get('wiener_iters', 0)
|
||||
if cac:
|
||||
zout = "Complex as Channels (CaC)"
|
||||
elif niters >= 0:
|
||||
zout = "Wiener filtering"
|
||||
else:
|
||||
zout = "Naive iSTFT from masking"
|
||||
self.add('___', zout, "Framework")
|
||||
self.add('cac', True, "Use Complex as Channels")
|
||||
if not self.get('cac', True):
|
||||
self.add('wiener_iters', 0, "Wiener filter iterations")
|
||||
self.flush("STFT")
|
||||
|
||||
def freq_branch(self):
|
||||
self.reset()
|
||||
def_ratio = None
|
||||
if self.get('multi_freqs') == []:
|
||||
def_ratio = []
|
||||
self.add('multi_freqs', def_ratio, "Ratios for frequency band splitting")
|
||||
if self.get('multi_freqs'):
|
||||
self.add('multi_freqs_depth', 2, "Layers to apply frequency splitting")
|
||||
self.add('freq_emb', 0.2, "Frequency embedding weight")
|
||||
if self.get('freq_emb'):
|
||||
indent = " "
|
||||
self.add('emb_scale', 10, "Scale", indent=indent)
|
||||
self.add('emb_smooth', True, "Smooth", indent=indent)
|
||||
self.flush("Frequency Branch")
|
||||
|
||||
def demucs(self):
|
||||
"""Logs the parameters for the original Demucs class."""
|
||||
self.log_type("Classic Waveform Demucs (Demucs)")
|
||||
self.log_param('sources', [], "Target source names")
|
||||
self.log_param('segment', 40, "Segment size", unit="s")
|
||||
# --- Structure & Channels ---
|
||||
self.structure(ch=64, depth=6, with_lstm=True)
|
||||
# --- Convolutions ---
|
||||
self.convolutions()
|
||||
|
||||
self.reset()
|
||||
self.add('gelu', True, "GeLU (not ReLU)")
|
||||
if self.get('rewrite', True):
|
||||
self.add('glu', True, "GLU in 1x1 rewrite (not ReLU)")
|
||||
self.flush("Activations")
|
||||
# --- Normalization ---
|
||||
self.normalization()
|
||||
# --- DConv Residual Branch ---
|
||||
self.dconv()
|
||||
# --- Pre/Post Processing ---
|
||||
self.reset()
|
||||
self.add('resample', True, "Use 2x resampling")
|
||||
self.add('normalize', True, "Normalize audio on-the-fly")
|
||||
self.flush("Processing")
|
||||
|
||||
def hdemucs(self):
|
||||
"""Logs the parameters for the HDemucs (Hybrid Spectrogram/Waveform) class."""
|
||||
self.log_type("Hybrid Demucs (Spectrogram + Waveform) (HDemucs)")
|
||||
self.log_param('sources', [], "Target source names")
|
||||
self.log_param('segment', 40, "Segment size", unit="s")
|
||||
# --- Structure & Channels ---
|
||||
self.structure(ch=48, depth=6, with_ch_tm=True)
|
||||
# --- STFT & Spectrogram ---
|
||||
self.stft()
|
||||
# --- Frequency Branch ---
|
||||
self.freq_branch()
|
||||
# --- Convolutions ---
|
||||
self.convolutions(advanced=True)
|
||||
# --- Normalization ---
|
||||
self.normalization()
|
||||
# --- DConv Residual Branch (defaults are different from Demucs) ---
|
||||
self.dconv()
|
||||
|
||||
def htdemucs(self):
|
||||
"""Logs the parameters for the HTDemucs (Hybrid Transformer) class."""
|
||||
self.log_type("Hybrid Transformer Demucs (HTDemucs)")
|
||||
self.log_param('sources', [], "Target source names")
|
||||
self.log_param('segment', 10, "Segment size", unit="s")
|
||||
# --- Structure & Channels (defaults are different from HDemucs) ---
|
||||
self.structure(ch=48, depth=4)
|
||||
# --- STFT & Spectrogram ---
|
||||
self.stft()
|
||||
# --- Frequency Branch ---
|
||||
self.freq_branch()
|
||||
# --- Convolutions ---
|
||||
self.convolutions(advanced=True)
|
||||
# --- Normalization ---
|
||||
self.normalization()
|
||||
# --- DConv (defaults are different) ---
|
||||
self.dconv(full=False)
|
||||
# --- Transformer Block ---
|
||||
if self.get('t_layers', 5) > 0:
|
||||
self.extra_indent = " "
|
||||
# --- Main Transformer ---
|
||||
self.reset()
|
||||
if self.get('bottom_channels', 0):
|
||||
self.add('bottom_channels', 0, "Channels forced to")
|
||||
self.add('t_hidden_scale', 4.0, "Hidden scale")
|
||||
self.add('t_layers', 5, "Number of transformer layers")
|
||||
self.add('t_heads', 8, "Number of attention heads")
|
||||
self.add('t_dropout', 0.0, "Dropout")
|
||||
self.flush("Transformer")
|
||||
# --- Positional Embeddings ---
|
||||
self.reset()
|
||||
self.add('t_emb', 'sin', "Type")
|
||||
self.add('t_weight_pos_embed', 1.0, "Weight", can_skip=True)
|
||||
t_emb = self.get('t_emb', 'sin')
|
||||
if t_emb == 'scaled':
|
||||
self.add('t_max_positions', 10000, "Max positions")
|
||||
elif t_emb == 'sin':
|
||||
self.add('t_max_period', 10000.0, "Max period")
|
||||
self.add('t_sin_random_shift', 0, "Random shift", can_skip=True)
|
||||
elif t_emb == 'cape':
|
||||
self.add('t_cape_mean_normalize', True, "Cape normalize")
|
||||
self.add('t_cape_glob_loc_scale', [5000.0, 1.0, 1.4], "Cape params")
|
||||
if self.get('t_cape_augment', True):
|
||||
rlogger.warning("t_cape_augment is True in loaded model, should be False for inference.")
|
||||
self.flush("Positional Embeddings", is_sub=True)
|
||||
# --- Transformer Normalization ---
|
||||
self.reset()
|
||||
self.add('t_norm_first', True, "Before attention/FFN")
|
||||
self.add('t_norm_in', True, "Before pos. embedding")
|
||||
if self.get('t_norm_in', True):
|
||||
self.add('t_norm_in_group', False, "On all timesteps")
|
||||
self.add('t_group_norm', False, "Of encoder on all timesteps")
|
||||
self.add('t_norm_out', True, "GroupNorm at end of layers")
|
||||
self.flush("Normalization", is_sub=True)
|
||||
# --- Transformer Misc ---
|
||||
self.reset()
|
||||
self.add('t_cross_first', False, "Cross-attention is the first layer")
|
||||
self.add('t_layer_scale', True, "Layer scale")
|
||||
self.add('t_gelu', True, "GeLU (not ReLU)")
|
||||
self.flush("Various", is_sub=True)
|
||||
# --- Sparsity ---
|
||||
# Log sparsity details only if sparse attention is enabled
|
||||
self.reset()
|
||||
is_sparse = self.get('t_sparse_self_attn', False)
|
||||
self.add('t_sparse_self_attn', False, "Use sparse self-attention")
|
||||
if is_sparse:
|
||||
self.add('t_sparse_cross_attn', False, "Sparse cross-attention")
|
||||
self.add('t_auto_sparsity', False, "Automatic sparsity")
|
||||
auto_sparsity = self.get('t_auto_sparsity', False)
|
||||
if not auto_sparsity:
|
||||
self.add('t_mask_type', 'diag', "Masking pattern")
|
||||
self.add('t_mask_random_seed', 42, "Mask seed")
|
||||
mask_t = self.get('t_mask_type', 'diag')
|
||||
if 'diag' in mask_t:
|
||||
self.add('t_sparse_attn_window', 500, "Window size")
|
||||
if 'global' in mask_t:
|
||||
self.add('t_global_window', 100, "Window size")
|
||||
if 'random' in mask_t:
|
||||
self.add('t_sparsity', 0.95, "Sparsity for random mask", unit="%")
|
||||
self.flush("Sparsity", is_sub=True)
|
||||
# Training only
|
||||
# self.add(logger, kwargs, 't_weight_decay', 0.0, "Weight decay", extra=False)
|
||||
# self.add(logger, kwargs, 't_lr', None, "Learning rate", extra=False)
|
||||
# self.add(logger, kwargs, 't_cape_augment', True, "Learning rate", extra=False)
|
||||
# self.add(logger, kwargs, 'rescale', 0.1, "Rescale trick", extra=False)
|
||||
self.extra_indent = ""
|
||||
@@ -8,15 +8,18 @@ import importlib
|
||||
import json
|
||||
import logging
|
||||
from safetensors import safe_open
|
||||
from seconohe.logger import get_debug_level
|
||||
from .MDX_Net import MDX_Net
|
||||
from ..utils.misc import NODES_NAME, json_object_hook
|
||||
from .. import NODES_NAME
|
||||
from ..utils.misc import json_object_hook
|
||||
# Demucs class imports
|
||||
from .demucs_api import BagOfModels
|
||||
from .demucs_log_helper import DemucsModelInfo
|
||||
|
||||
logger = logging.getLogger(f"{NODES_NAME}.get_model")
|
||||
|
||||
|
||||
def get_metadata(file_path, d):
|
||||
def get_metadata(file_path, d=None):
|
||||
""" Read the metadata from a safetensors file """
|
||||
logger.debug(f"Reading metadata from {file_path}")
|
||||
metadata = {}
|
||||
@@ -25,6 +28,8 @@ def get_metadata(file_path, d):
|
||||
if not metadata:
|
||||
raise ValueError(f"Could not read metadata from safetensors file: {file_path}")
|
||||
|
||||
if d is None:
|
||||
return metadata
|
||||
# Is this a child model?
|
||||
parent = d.get('parent')
|
||||
if parent:
|
||||
@@ -57,7 +62,7 @@ def get_hyperparameter(metadata, parameter, as_type, default=None, warn_diff=Tru
|
||||
def get_mdx_model(d):
|
||||
""" Create an MDX_Net object with the specified parameters """
|
||||
# Check the file is consistent we our data base
|
||||
metadata = get_metadata(d['model_path'])
|
||||
metadata = get_metadata(d['model_path'], d)
|
||||
dim_f = get_hyperparameter(metadata, 'mdx_dim_f_set', "int", d['mdx_dim_f_set'])
|
||||
channels = get_hyperparameter(metadata, 'channels', "int", d['channels'])
|
||||
stages = get_hyperparameter(metadata, 'stages', "int", d['stages'])
|
||||
@@ -102,8 +107,11 @@ def get_demucs_model(d):
|
||||
logger.debug(f" - Redirecting class: {class_module} -> {local_class_module}")
|
||||
module = importlib.import_module(local_class_module)
|
||||
klass = getattr(module, class_name)
|
||||
instance = klass(*args, **kwargs)
|
||||
instance._signature = sig
|
||||
instance._metadata = model_meta
|
||||
|
||||
sub_models.append({'sig': sig, 'model': klass(*args, **kwargs)})
|
||||
sub_models.append({'sig': sig, 'model': instance})
|
||||
|
||||
# Create a mapping from signature to model instance
|
||||
model_map = {m['sig']: m['model'] for m in sub_models}
|
||||
@@ -113,17 +121,37 @@ def get_demucs_model(d):
|
||||
segment = float(metadata.get('segment', '0'))
|
||||
|
||||
if not is_bag:
|
||||
weights = None
|
||||
final_model = ordered_models[0]
|
||||
final_model.signatures = signatures
|
||||
else:
|
||||
logger.debug("Rebuilding BagOfModels container...")
|
||||
weights = json.loads(metadata.get('weights', 'null'))
|
||||
|
||||
final_model = BagOfModels(ordered_models, weights=weights, segment=segment)
|
||||
final_model.signatures = signatures
|
||||
|
||||
final_model.config_segment = segment
|
||||
|
||||
debug_level = get_debug_level(logger)
|
||||
if debug_level >= 1:
|
||||
# Show some information of the resulting model
|
||||
logger.debug("Model information:")
|
||||
logger.debug(f"- Total models {len(ordered_models)}")
|
||||
if weights is not None and len(weights) != len(ordered_models):
|
||||
raise ValueError(f"Invalid weights for {len(ordered_models)} models: {weights}")
|
||||
for n, m in enumerate(ordered_models):
|
||||
tp = m.__class__.__name__
|
||||
w = weights[n] if weights else None
|
||||
if n == 0:
|
||||
ref_sources = m.sources
|
||||
else:
|
||||
if m.sources != ref_sources:
|
||||
logger.error("The sub-model outputs doesn't match")
|
||||
if w and len(m.sources) != len(w):
|
||||
raise ValueError(f"Invalid {w} weights for {m.sources} sources")
|
||||
num = -1 if len(ordered_models) == 1 else n
|
||||
DemucsModelInfo(num, tp, m._metadata['kwargs'], logger.debug, w, extra=debug_level > 1, sig=m._signature)
|
||||
|
||||
return final_model
|
||||
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
# Short-Time Fourier Transform (STFT).
|
||||
import logging
|
||||
import numpy as np
|
||||
from seconohe.torch import model_to_target
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
# Local imports
|
||||
from ..utils.misc import NODES_NAME
|
||||
from ..utils.torch import model_to_target
|
||||
from .. import NODES_NAME
|
||||
|
||||
logger = logging.getLogger(f"{NODES_NAME}.stft")
|
||||
|
||||
@@ -127,7 +127,7 @@ def stft_chunk_process(waveform, d, model_run, device, segment_size=256, hop_len
|
||||
total_chunks = 1 + (mixture.shape[1] - chunk_size + step - 1) // step
|
||||
logger.info(f"⚙️ Processing {total_chunks} chunks...")
|
||||
|
||||
with model_to_target(model_run):
|
||||
with model_to_target(logger, model_run):
|
||||
for i in tqdm(range(0, mixture.shape[1] - chunk_size + 1, step)):
|
||||
start = i
|
||||
end = i + chunk_size
|
||||
@@ -4,15 +4,16 @@
|
||||
# Project: ComfyUI-AudioSeparation
|
||||
import os
|
||||
from typing import Dict
|
||||
from seconohe.comfy_node_action import send_node_action
|
||||
from seconohe.torch import get_torch_device_options, get_canonical_device
|
||||
# ComfyUI imports
|
||||
import folder_paths # ComfyUI's way to access model paths
|
||||
# Local imports
|
||||
from .source.utils.logger import main_logger
|
||||
from .source.utils.load_audio import audio_get_channels, force_stereo, force_sample_rate
|
||||
from .source.utils.torch import get_torch_device_options, get_canonical_device
|
||||
from .source.utils.comfy_node_action import send_node_action
|
||||
from .source.db.models_db import ModelsDB
|
||||
from .source.inference.demixer import get_demixer
|
||||
# We are the main source, so we use the main_logger
|
||||
from . import main_logger
|
||||
from .utils.load_audio import audio_get_channels, force_stereo, force_sample_rate
|
||||
from .db.models_db import ModelsDB
|
||||
from .inference.demixer import get_demixer
|
||||
|
||||
DEF_MODEL = 'Kim_Vocal_2.safetensors'
|
||||
DEF_ENTRY = 'Default'
|
||||
@@ -94,7 +95,7 @@ class AudioSeparateVocals:
|
||||
# Handle a change in the icon of the model name
|
||||
if model_path is None:
|
||||
# Was downloaded
|
||||
send_node_action("change_widget", "model", model_data['indicator'] + model_data['filtered_name'])
|
||||
send_node_action(logger, "change_widget", "model", model_data['indicator'] + model_data['filtered_name'])
|
||||
|
||||
# Match channels and S/R
|
||||
waveform = input_sound['waveform']
|
||||
@@ -220,7 +221,7 @@ class AudioSeparateDemucs(AudioSeparateVocals):
|
||||
# Handle a change in the icon of the model name
|
||||
if model_path is None:
|
||||
# Was downloaded
|
||||
send_node_action("change_widget", "model", model_data['indicator'] + model_data['filtered_name'])
|
||||
send_node_action(logger, "change_widget", "model", model_data['indicator'] + model_data['filtered_name'])
|
||||
|
||||
# Match channels and S/R
|
||||
waveform = input_sound['waveform']
|
||||
@@ -8,7 +8,7 @@
|
||||
import logging
|
||||
import torch
|
||||
import torchaudio
|
||||
from .misc import NODES_NAME
|
||||
from .. import NODES_NAME
|
||||
|
||||
logger = logging.getLogger(f"{NODES_NAME}.load_audio")
|
||||
|
||||
@@ -8,7 +8,7 @@ import importlib
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from .misc import NODES_NAME
|
||||
from .. import NODES_NAME
|
||||
|
||||
logger = logging.getLogger(f"{NODES_NAME}.load_class")
|
||||
|
||||
@@ -11,7 +11,7 @@ try:
|
||||
with_onnx = True
|
||||
except Exception:
|
||||
with_onnx = False
|
||||
from .misc import NODES_NAME
|
||||
from .. import NODES_NAME
|
||||
|
||||
logger = logging.getLogger(f"{NODES_NAME}.load_onnx")
|
||||
|
||||
@@ -7,7 +7,7 @@
|
||||
# Original code from Gemini 2.5 Pro
|
||||
import logging
|
||||
from safetensors.torch import load_file
|
||||
from .misc import NODES_NAME
|
||||
from .. import NODES_NAME
|
||||
|
||||
logger = logging.getLogger(f"{NODES_NAME}.load_safetensors")
|
||||
|
||||
@@ -2,17 +2,10 @@
|
||||
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
# License: GPLv3
|
||||
# Project: ComfyUI-AudioSeparation
|
||||
import argparse
|
||||
from fractions import Fraction
|
||||
import json
|
||||
import logging
|
||||
|
||||
NODES_NAME = "AudioSeparation"
|
||||
NODES_DEBUG_VAR = NODES_NAME.upper() + "_NODES_DEBUG"
|
||||
|
||||
|
||||
def debugl(logger, level, msg):
|
||||
if logger.getEffectiveLevel() <= logging.DEBUG - (level - 1):
|
||||
logger.debug(msg)
|
||||
from .. import __version__, __copyright__, __license__, __author__
|
||||
|
||||
|
||||
def cli_add_verbose(parser):
|
||||
@@ -20,6 +13,29 @@ def cli_add_verbose(parser):
|
||||
help="Enable verbose output to see details of the process.")
|
||||
|
||||
|
||||
class PrintVersionAction(argparse.Action):
|
||||
def __init__(self, option_strings, dest, nargs=None, **kwargs):
|
||||
super().__init__(option_strings, dest, nargs=0, **kwargs)
|
||||
|
||||
def __call__(self, parser, namespace, values, option_string=None):
|
||||
# Format the version information
|
||||
version_info = f"""{parser.prog} (Audio Separation) {__version__}
|
||||
{__copyright__}
|
||||
{__license__}
|
||||
This is free software: you are free to change and redistribute it.
|
||||
There is NO WARRANTY, to the extent permitted by law.
|
||||
|
||||
Written by {__author__}"""
|
||||
print(version_info)
|
||||
# Exit the parser
|
||||
parser.exit()
|
||||
|
||||
|
||||
def cli_add_version(parser, prog_name):
|
||||
parser.add_argument('-V', '--version', help="Show version and copyright information and exit",
|
||||
action=PrintVersionAction)
|
||||
|
||||
|
||||
class FractionEncoder(json.JSONEncoder):
|
||||
def default(self, obj):
|
||||
if isinstance(obj, Fraction):
|
||||
@@ -8,7 +8,7 @@
|
||||
import logging
|
||||
import os
|
||||
import torchaudio
|
||||
from .misc import NODES_NAME
|
||||
from .. import NODES_NAME
|
||||
|
||||
logger = logging.getLogger(f"{NODES_NAME}.save_audio")
|
||||
|
||||
Regular → Executable
+8
-5
@@ -1,3 +1,4 @@
|
||||
#!/usr/bin/env python3
|
||||
# Copyright (c) 2025 Salvador E. Tropea
|
||||
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
# License: GPLv3
|
||||
@@ -11,14 +12,15 @@ import argparse
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from seconohe.logger import logger_set_standalone
|
||||
import subprocess
|
||||
import sys
|
||||
# Local imports
|
||||
import bootstrap # noqa: F401
|
||||
from source.db.hash import get_hash
|
||||
from source.db.models_db import load_known_models, save_known_models, get_db_filename
|
||||
from source.utils.logger import main_logger, logger_set_standalone
|
||||
from source.utils.misc import cli_add_verbose
|
||||
from src.nodes import main_logger
|
||||
from src.nodes.db.hash import get_hash
|
||||
from src.nodes.db.models_db import load_known_models, save_known_models, get_db_filename
|
||||
from src.nodes.utils.misc import cli_add_verbose, cli_add_version
|
||||
|
||||
|
||||
def parse_converter_output(output):
|
||||
@@ -37,7 +39,7 @@ def parse_converter_output(output):
|
||||
|
||||
|
||||
def main(args):
|
||||
logger_set_standalone(args)
|
||||
logger_set_standalone(main_logger, args)
|
||||
# 1. Load the JSON metadata file
|
||||
model_db = load_known_models(args.json_file)
|
||||
if model_db is None:
|
||||
@@ -200,6 +202,7 @@ if __name__ == "__main__":
|
||||
parser.add_argument('--model_location', type=str, default='source/inference/MDX_Net.py:MDX_Net',
|
||||
help="Python path to the model class.")
|
||||
cli_add_verbose(parser)
|
||||
cli_add_version(parser, __name__)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
|
||||
Executable
+80
@@ -0,0 +1,80 @@
|
||||
#!/usr/bin/env python3
|
||||
import re
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
# --- Configuration ---
|
||||
# The path to your pyproject.toml file, relative to the project root
|
||||
PYPROJECT_PATH = Path("pyproject.toml")
|
||||
# The path to the Python file containing the __version__ string
|
||||
SOURCE_VERSION_PATH = Path("src/nodes/__init__.py")
|
||||
# --- End Configuration ---
|
||||
|
||||
|
||||
def get_version_from_pyproject(file_path: Path) -> str | None:
|
||||
"""Extracts the version string from a pyproject.toml file."""
|
||||
try:
|
||||
content = file_path.read_text()
|
||||
# A simple regex to find `version = "..."` under the `[project]` table
|
||||
match = re.search(r'^version\s*=\s*"(.*?)"', content, re.MULTILINE)
|
||||
if match:
|
||||
return match.group(1)
|
||||
except FileNotFoundError:
|
||||
print(f"Error: {file_path} not found.", file=sys.stderr)
|
||||
except Exception as e:
|
||||
print(f"Error reading or parsing {file_path}: {e}", file=sys.stderr)
|
||||
return None
|
||||
|
||||
|
||||
def get_version_from_source(file_path: Path) -> str | None:
|
||||
"""Extracts the __version__ string from a Python source file."""
|
||||
try:
|
||||
content = file_path.read_text()
|
||||
# A simple regex to find `__version__ = "..."`
|
||||
match = re.search(r'^__version__\s*=\s*"(.*?)"', content, re.MULTILINE)
|
||||
if match:
|
||||
return match.group(1)
|
||||
except FileNotFoundError:
|
||||
print(f"Error: {file_path} not found.", file=sys.stderr)
|
||||
except Exception as e:
|
||||
print(f"Error reading or parsing {file_path}: {e}", file=sys.stderr)
|
||||
return None
|
||||
|
||||
|
||||
def main() -> int:
|
||||
"""
|
||||
Compares version strings from pyproject.toml and the source code.
|
||||
Exits with a non-zero status code if they do not match.
|
||||
"""
|
||||
print("--- Checking version consistency ---")
|
||||
|
||||
# Get versions
|
||||
pyproject_version = get_version_from_pyproject(PYPROJECT_PATH)
|
||||
source_version = get_version_from_source(SOURCE_VERSION_PATH)
|
||||
|
||||
# Validate that we found both
|
||||
if not pyproject_version:
|
||||
print(f"Error: Could not find version in {PYPROJECT_PATH}", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
if not source_version:
|
||||
print(f"Error: Could not find `__version__` in {SOURCE_VERSION_PATH}", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
print(f"Version in {PYPROJECT_PATH}: {pyproject_version}")
|
||||
print(f"Version in {SOURCE_VERSION_PATH}: {source_version}")
|
||||
|
||||
# Compare and exit
|
||||
if pyproject_version == source_version:
|
||||
print("✅ Versions are consistent.")
|
||||
return 0
|
||||
else:
|
||||
print("\n❌ Error: Version mismatch!", file=sys.stderr)
|
||||
print(f" pyproject.toml has version '{pyproject_version}'", file=sys.stderr)
|
||||
print(f" {SOURCE_VERSION_PATH} has version '{source_version}'", file=sys.stderr)
|
||||
print(" Please ensure both versions are identical.", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
Regular → Executable
+12
-9
@@ -1,3 +1,4 @@
|
||||
#!/usr/bin/env python3
|
||||
# Copyright (c) 2025 Salvador E. Tropea
|
||||
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
# License: GPLv3
|
||||
@@ -7,16 +8,17 @@
|
||||
# Run it using: python tool/demix.py -m HASH AUDIO
|
||||
import argparse
|
||||
import os
|
||||
from seconohe.logger import logger_set_standalone
|
||||
from seconohe.torch import get_torch_device_options, get_canonical_device
|
||||
import sys
|
||||
# Local imports
|
||||
import bootstrap # noqa: F401
|
||||
from source.db.models_db import ModelsDB, cli_add_models_and_db
|
||||
from source.inference.demixer import get_demixer
|
||||
from source.utils.load_audio import load_audio
|
||||
from source.utils.logger import main_logger, logger_set_standalone
|
||||
from source.utils.misc import cli_add_verbose
|
||||
from source.utils.save_audio import save_audio
|
||||
from source.utils.torch import get_torch_device_options, get_canonical_device
|
||||
from src.nodes import main_logger
|
||||
from src.nodes.db.models_db import ModelsDB, cli_add_models_and_db
|
||||
from src.nodes.inference.demixer import get_demixer
|
||||
from src.nodes.utils.load_audio import load_audio
|
||||
from src.nodes.utils.misc import cli_add_verbose, cli_add_version
|
||||
from src.nodes.utils.save_audio import save_audio
|
||||
|
||||
BANNER = "🎵 MDX-Net Audio Separation Tool 🎵"
|
||||
|
||||
@@ -47,7 +49,7 @@ def demix(d, args):
|
||||
|
||||
if args.save_complement:
|
||||
for wav in wavs[1:]:
|
||||
if wav is None:
|
||||
if wav is None or not wav['generated']:
|
||||
continue
|
||||
out_path = f"{args.out_base or base}_{wav['stem']}{out_ext}"
|
||||
save_audio(wav['waveform'], wav['sample_rate'], out_path, out_format)
|
||||
@@ -84,10 +86,11 @@ if __name__ == "__main__":
|
||||
help="How many audio segments to process at once. 0 uses the value from the model")
|
||||
parser.add_argument('-l', '--list', action='store_true', help="Show available models")
|
||||
cli_add_verbose(parser)
|
||||
cli_add_version(parser, __name__)
|
||||
|
||||
parser.set_defaults(save_main=True)
|
||||
args = parser.parse_args()
|
||||
logger_set_standalone(args)
|
||||
logger_set_standalone(main_logger, args)
|
||||
main_logger.info(BANNER)
|
||||
|
||||
# Sanity check
|
||||
|
||||
Regular → Executable
+52
-10
@@ -1,3 +1,4 @@
|
||||
#!/usr/bin/env python3
|
||||
# Copyright (c) 2025 Salvador E. Tropea
|
||||
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
# License: GPLv3
|
||||
@@ -13,6 +14,7 @@ import argparse
|
||||
import json
|
||||
from pathlib import Path
|
||||
from safetensors.torch import save_file
|
||||
from seconohe.logger import debugl, logger_set_standalone
|
||||
import sys
|
||||
import torch
|
||||
from typing import Dict
|
||||
@@ -25,17 +27,21 @@ try:
|
||||
except Exception:
|
||||
with_demuc_lib = False
|
||||
import bootstrap # noqa: F401
|
||||
from source.utils.misc import cli_add_verbose, FractionEncoder
|
||||
from source.utils.logger import main_logger, logger_set_standalone
|
||||
from source.db.models_db import cli_add_db, get_download_url, ModelsDB
|
||||
from source.db.hash import get_hash
|
||||
import source.inference.Demucs as local_demucs_module
|
||||
import source.inference.HDemucs as local_hdemucs_module
|
||||
import source.inference.HTDemucs as local_htdemucs_module
|
||||
from src.nodes import main_logger
|
||||
from src.nodes.utils.misc import cli_add_verbose, FractionEncoder, cli_add_version
|
||||
from src.nodes.db.models_db import cli_add_db, get_download_url, ModelsDB
|
||||
from src.nodes.db.hash import get_hash
|
||||
import src.nodes.inference.Demucs as local_demucs_module
|
||||
import src.nodes.inference.HDemucs as local_hdemucs_module
|
||||
import src.nodes.inference.HTDemucs as local_htdemucs_module
|
||||
MODULES_MAP = {'demucs': local_demucs_module,
|
||||
'demucs.demucs': local_demucs_module,
|
||||
'demucs.hdemucs': local_hdemucs_module,
|
||||
'demucs.htdemucs': local_htdemucs_module}
|
||||
MAP = {'freq_encoder': 'encoder',
|
||||
'freq_decoder': 'decoder',
|
||||
'time_encoder': 'tencoder',
|
||||
'time_decoder': 'tdecoder'}
|
||||
logger = main_logger
|
||||
|
||||
|
||||
@@ -61,6 +67,36 @@ def remap_module(modules_map):
|
||||
del sys.modules[old_name]
|
||||
|
||||
|
||||
def solve_simple_pt(yaml_data, pkg):
|
||||
""" PyTorch Audio lib has some raw models, we store the metadata in the YAML """
|
||||
klass = yaml_data.get('klass')
|
||||
if klass is None:
|
||||
main_logger.error("No `klass` in YAML")
|
||||
sys.exit(4)
|
||||
if klass == 'Demucs':
|
||||
klass = local_demucs_module.Demucs
|
||||
elif klass == 'HDemucs':
|
||||
klass = local_hdemucs_module.HDemucs
|
||||
elif klass == 'HTDemucs':
|
||||
klass = local_htdemucs_module.HTDemucs
|
||||
else:
|
||||
main_logger.error("Unknown model `klass` {klass}")
|
||||
sys.exit(4)
|
||||
args = yaml_data.get('args', {})
|
||||
kwargs = yaml_data.get('kwargs', {})
|
||||
|
||||
# For PyTorch Audio model (very old code?)
|
||||
new_dict = {}
|
||||
for k, v in pkg.items():
|
||||
parts = k.split('.')
|
||||
gr = parts[0]
|
||||
if gr in MAP:
|
||||
k = MAP[gr] + '.' + '.'.join(parts[1:])
|
||||
new_dict[k] = v
|
||||
|
||||
return klass, args, kwargs, new_dict
|
||||
|
||||
|
||||
def convert_demucs_model(yaml_path_str: str, model_paths: list[str], output_path_str: str, all_metadata, data: Dict):
|
||||
"""
|
||||
Loads an original Demucs model bag, extracts weights and all necessary metadata
|
||||
@@ -111,7 +147,11 @@ def convert_demucs_model(yaml_path_str: str, model_paths: list[str], output_path
|
||||
# This is the only "unsafe" part, loading the original pickle file
|
||||
pkg = torch.load(path, map_location='cpu', weights_only=False)
|
||||
|
||||
klass, args, kwargs, state = pkg["klass"], pkg["args"], pkg["kwargs"], pkg["state"]
|
||||
debugl(logger, 2, f"PyTorch data type is {type(pkg)}")
|
||||
if isinstance(pkg, dict) and 'klass' in pkg:
|
||||
klass, args, kwargs, state = pkg["klass"], pkg["args"], pkg["kwargs"], pkg["state"]
|
||||
else:
|
||||
klass, args, kwargs, state = solve_simple_pt(yaml_data, pkg)
|
||||
main_logger.info(f" - Model class: {klass.__module__}.{klass.__name__}")
|
||||
|
||||
# Dequantize if necessary by letting the original code handle it
|
||||
@@ -162,6 +202,7 @@ def convert_demucs_model(yaml_path_str: str, model_paths: list[str], output_path
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser(description="Convert Demucs .th models to a single .safetensors file.")
|
||||
|
||||
parser.add_argument('--yaml', required=True, type=str, help="Path to the Demucs .yaml file.")
|
||||
parser.add_argument('--models', nargs='*', default=None, type=str,
|
||||
help="Optional. Paths to .th files. If not provided, assumes they are "
|
||||
@@ -171,9 +212,10 @@ if __name__ == '__main__':
|
||||
"it's saved next to the YAML with the same name.")
|
||||
cli_add_db(parser)
|
||||
cli_add_verbose(parser)
|
||||
cli_add_version(parser, __name__)
|
||||
|
||||
args = parser.parse_args()
|
||||
logger_set_standalone(args)
|
||||
logger_set_standalone(main_logger, args)
|
||||
main_logger.info("⚙️ PyTorch Demucs to Safetensors converter\n")
|
||||
|
||||
# Get information about the YAML file in our database
|
||||
@@ -205,7 +247,7 @@ if __name__ == '__main__':
|
||||
|
||||
model_files = []
|
||||
for sig in set(signatures): # Use set to avoid redundant searches
|
||||
found = list(yaml_dir.glob(f'{sig}*.th'))
|
||||
found = list(yaml_dir.glob(f'{sig}*.th')) + list(yaml_dir.glob(f'{sig}*.pt'))
|
||||
if not found:
|
||||
raise FileNotFoundError(f"Could not automatically find a model file for signature '{sig}' in {yaml_dir}")
|
||||
if len(found) > 1:
|
||||
|
||||
Regular → Executable
+1
@@ -1,3 +1,4 @@
|
||||
#!/usr/bin/env python3
|
||||
# Copyright (c) 2025 Salvador E. Tropea
|
||||
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
# License: GPLv3
|
||||
|
||||
Regular → Executable
+9
-6
@@ -1,3 +1,4 @@
|
||||
#!/usr/bin/env python3
|
||||
# Copyright (c) 2025 Salvador E. Tropea
|
||||
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
# License: GPLv3
|
||||
@@ -12,16 +13,17 @@ import json
|
||||
import numpy as np
|
||||
import onnx
|
||||
import os
|
||||
from safetensors.torch import save_file
|
||||
from seconohe.logger import logger_set_standalone
|
||||
import sys
|
||||
import torch
|
||||
from torch import nn
|
||||
# Local imports
|
||||
import bootstrap # noqa: F401
|
||||
from safetensors.torch import save_file
|
||||
from source.utils.load_class import import_model_class
|
||||
from source.utils.logger import main_logger, logger_set_standalone
|
||||
from source.utils.misc import cli_add_verbose
|
||||
from source.db.models_db import get_download_url
|
||||
from src.nodes import main_logger
|
||||
from src.nodes.utils.load_class import import_model_class
|
||||
from src.nodes.utils.misc import cli_add_verbose, cli_add_version
|
||||
from src.nodes.db.models_db import get_download_url
|
||||
|
||||
|
||||
class OnnxGraph:
|
||||
@@ -224,9 +226,10 @@ if __name__ == "__main__":
|
||||
"(default: source/inference/MDX_Net.py:MDX_Net)")
|
||||
parser.add_argument('-j', '--metadata', type=str, help="Metadata to include in the hyperparameters, JSON format")
|
||||
cli_add_verbose(parser)
|
||||
cli_add_version(parser, __name__)
|
||||
|
||||
args = parser.parse_args()
|
||||
logger_set_standalone(args)
|
||||
logger_set_standalone(main_logger, args)
|
||||
if args.metadata is not None:
|
||||
try:
|
||||
args.metadata = json.loads(args.metadata)
|
||||
|
||||
Regular → Executable
+8
-5
@@ -1,3 +1,4 @@
|
||||
#!/usr/bin/env python3
|
||||
# Copyright (c) 2025 Salvador E. Tropea
|
||||
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
# License: GPLv3
|
||||
@@ -6,13 +7,14 @@
|
||||
# Tool to show PyTorch class i.e:
|
||||
# python tool/show_class.py -m source/inference/MDX_Net.py:MDX_Net
|
||||
import argparse
|
||||
from seconohe.logger import logger_set_standalone
|
||||
import sys
|
||||
from torch import nn
|
||||
# Local imports
|
||||
import bootstrap # noqa: F401
|
||||
from source.utils.logger import main_logger, logger_set_standalone
|
||||
from source.utils.load_class import import_model_class
|
||||
from source.utils.misc import cli_add_verbose
|
||||
from src.nodes import main_logger
|
||||
from src.nodes.utils.load_class import import_model_class
|
||||
from src.nodes.utils.misc import cli_add_verbose, cli_add_version
|
||||
|
||||
|
||||
def load_class(args):
|
||||
@@ -121,16 +123,17 @@ if __name__ == "__main__":
|
||||
parser.add_argument('-n', '--num_stages', type=int, default=5,
|
||||
choices=[2, 3, 4, 5, 6, 7], # Restrict to known valid values
|
||||
help="The number of U-Net stages in the model. (choices: 2 to 7, default: 5)")
|
||||
cli_add_verbose(parser)
|
||||
parser.add_argument('-o', '--export_onnx', type=str, default=None,
|
||||
help="Path for the optional output .onnx file.\n"
|
||||
"Only the structure is exported")
|
||||
parser.add_argument('-k', '--keys', action='store_true', help="Print the keys for the state_dict.")
|
||||
parser.add_argument('-C', '--compact', action='store_true', help="Print a compact representation.")
|
||||
parser.add_argument('-S', '--no_show', action='store_false', help="Don't print the structure.")
|
||||
cli_add_verbose(parser)
|
||||
cli_add_version(parser, __name__)
|
||||
|
||||
args = parser.parse_args()
|
||||
logger_set_standalone(args)
|
||||
logger_set_standalone(main_logger, args)
|
||||
model = load_class(args)
|
||||
if args.no_show:
|
||||
show(model)
|
||||
|
||||
Regular → Executable
+8
-5
@@ -1,3 +1,4 @@
|
||||
#!/usr/bin/env python3
|
||||
# Copyright (c) 2025 Salvador E. Tropea
|
||||
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
# License: GPLv3
|
||||
@@ -8,13 +9,14 @@
|
||||
# Run it using: python tool/show_db.py
|
||||
import argparse
|
||||
import pprint
|
||||
from seconohe.logger import logger_set_standalone
|
||||
import sys
|
||||
# Local imports
|
||||
import bootstrap # noqa: F401
|
||||
from source.db.hash_dir import hash_dir
|
||||
from source.db.models_db import load_known_models, cli_add_models_and_db, save_known_models, get_models
|
||||
from source.utils.logger import main_logger, logger_set_standalone
|
||||
from source.utils.misc import cli_add_verbose
|
||||
from src.nodes import main_logger
|
||||
from src.nodes.db.hash_dir import hash_dir
|
||||
from src.nodes.db.models_db import load_known_models, cli_add_models_and_db, save_known_models, get_models
|
||||
from src.nodes.utils.misc import cli_add_verbose, cli_add_version
|
||||
|
||||
|
||||
# Do nothing, you can apply some change here
|
||||
@@ -63,7 +65,7 @@ def apply_process(model_db):
|
||||
|
||||
|
||||
def main(args):
|
||||
logger_set_standalone(args)
|
||||
logger_set_standalone(main_logger, args)
|
||||
# 1. Load the JSON metadata file
|
||||
model_db = load_known_models(args.json_file)
|
||||
if model_db is None:
|
||||
@@ -128,6 +130,7 @@ if __name__ == "__main__":
|
||||
|
||||
# --- Control Arguments ---
|
||||
cli_add_verbose(parser)
|
||||
cli_add_version(parser, __name__)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
|
||||
Regular → Executable
+23
-5
@@ -1,3 +1,4 @@
|
||||
#!/usr/bin/env python3
|
||||
# Copyright (c) 2025 Salvador E. Tropea
|
||||
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
# License: GPLv3
|
||||
@@ -8,11 +9,13 @@
|
||||
import argparse
|
||||
import json
|
||||
from pprint import pprint
|
||||
from seconohe.logger import logger_set_standalone
|
||||
# Local imports
|
||||
import bootstrap # noqa: F401
|
||||
from source.utils.logger import main_logger, logger_set_standalone
|
||||
from source.utils.misc import cli_add_verbose, json_object_hook
|
||||
from source.inference.get_model import get_metadata
|
||||
from src.nodes import main_logger
|
||||
from src.nodes.utils.misc import cli_add_verbose, json_object_hook, cli_add_version
|
||||
from src.nodes.inference.get_model import get_metadata
|
||||
from src.nodes.inference.demucs_log_helper import DemucsModelInfo
|
||||
|
||||
|
||||
def show_metadata(args):
|
||||
@@ -23,13 +26,28 @@ def show_metadata(args):
|
||||
expanded = {k: json.loads(v, object_hook=json_object_hook) if v[0] in '{[' else v for k, v in d.items()}
|
||||
pprint(expanded)
|
||||
|
||||
sigs = expanded.get('signatures')
|
||||
if sigs:
|
||||
single = len(sigs) == 1
|
||||
print("Demucs model")
|
||||
if not single:
|
||||
print(f"Composed by {len(sigs)} submodels")
|
||||
weights = expanded.get('weights')
|
||||
# Demucs model
|
||||
for n, s in enumerate(sigs):
|
||||
d = expanded[s]
|
||||
num = -1 if single else n
|
||||
DemucsModelInfo(num, d['class_name'], d['kwargs'], print, weights[n] if weights else None, extra=args.verbose,
|
||||
sig=s)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="Shows the safetensors metadata",
|
||||
formatter_class=argparse.RawTextHelpFormatter)
|
||||
parser.add_argument('input_file', type=str, help="Path to the input ONNX model file.")
|
||||
parser.add_argument('input_file', type=str, help="Path to the input safetensors model file.")
|
||||
cli_add_verbose(parser)
|
||||
cli_add_version(parser, __name__)
|
||||
|
||||
args = parser.parse_args()
|
||||
logger_set_standalone(args)
|
||||
logger_set_standalone(main_logger, args)
|
||||
show_metadata(args)
|
||||
|
||||
Regular → Executable
+6
-3
@@ -1,3 +1,4 @@
|
||||
#!/usr/bin/env python3
|
||||
# Copyright (c) 2025 Salvador E. Tropea
|
||||
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
|
||||
# License: GPLv3
|
||||
@@ -10,11 +11,12 @@ import numpy as np
|
||||
import onnx
|
||||
from onnx import shape_inference # Import the shape inference module
|
||||
import torch
|
||||
from seconohe.logger import logger_set_standalone
|
||||
import sys
|
||||
# Local imports
|
||||
import bootstrap # noqa: F401
|
||||
from source.utils.logger import main_logger, logger_set_standalone
|
||||
from source.utils.misc import cli_add_verbose
|
||||
from src.nodes import main_logger
|
||||
from src.nodes.utils.misc import cli_add_verbose, cli_add_version
|
||||
|
||||
|
||||
def print_onnx_nodes_and_weights(onnx_model_path):
|
||||
@@ -202,9 +204,10 @@ if __name__ == "__main__":
|
||||
"Incompatible with -S")
|
||||
parser.add_argument('-S', '--no_show', action='store_false', help="Don't print the ONNX structure.")
|
||||
cli_add_verbose(parser)
|
||||
cli_add_version(parser, __name__)
|
||||
|
||||
args = parser.parse_args()
|
||||
logger_set_standalone(args)
|
||||
logger_set_standalone(main_logger, args)
|
||||
if args.run and not args.no_show:
|
||||
main_logger.error("-r can't be used when -S is specified")
|
||||
sys.exit(1)
|
||||
|
||||
Regular → Executable
+4
-3
@@ -8,10 +8,11 @@
|
||||
# python tool/uvr_hash.py model.onnx
|
||||
import argparse
|
||||
import os
|
||||
from seconohe.logger import logger_set_standalone
|
||||
# Local imports
|
||||
import bootstrap # noqa: F401
|
||||
from source.utils.logger import main_logger, logger_set_standalone
|
||||
from source.db.hash import get_hash
|
||||
from src.nodes import main_logger
|
||||
from src.nodes.db.hash import get_hash
|
||||
|
||||
|
||||
def main():
|
||||
@@ -36,7 +37,7 @@ def main():
|
||||
|
||||
args = parser.parse_args()
|
||||
args.verbose = 0
|
||||
logger_set_standalone(args)
|
||||
logger_set_standalone(main_logger, args)
|
||||
|
||||
# Process each file provided on the command line
|
||||
for filepath in args.files:
|
||||
|
||||
Reference in New Issue
Block a user