30 Commits
Author SHA1 Message Date
Salvador E. Tropea 621bd2784e Bumped version to 1.1.3 2026-02-11 11:03:35 -03:00
Salvador E. Tropea 4e56851264 [Fixed] Issues on Windows when auto-downloading
Fixes #4
2026-02-11 10:53:37 -03:00
Salvador E. Tropea 8a84f6b48e Bumped version to 1.1.2 2025-07-27 14:40:42 -03:00
Salvador E. Tropea 2737d7b190 [Added] Version information
- When registering the nodes
- To command line tools
2025-07-27 14:38:26 -03:00
Salvador E. Tropea 335af9fa4b [Added] Pre-commit script to check version consistency
ComfyUI registry fault, poor support of pyproject.toml options
2025-07-27 14:36:23 -03:00
Salvador E. Tropea accb0fd56c [DOCs][Added] README for the icons 2025-07-27 12:48:57 -03:00
Salvador E. Tropea 329a8826ad [CI/CD][Changed] To publish on tag (with semantic version) 2025-07-27 12:44:52 -03:00
Salvador E. Tropea d8f885e61a [DOCs] Updated installation and description
Also added Icon and Banner entries for ComfyUI manager
2025-07-27 12:40:20 -03:00
Salvador E. Tropea 82e0f78e21 Ignore sync script 2025-07-27 12:34:27 -03:00
Salvador E. Tropea 17de6ca853 [Torch] Moved to SeCoNoHe 2025-07-23 13:15:25 -03:00
Salvador E. Tropea 87aacaa903 [Tools] Adapted to SeCoNoHe 2025-07-23 10:44:11 -03:00
Salvador E. Tropea 4da4a7b498 [Logger] Now using get_debug_level and debugl from SeCoNoHe 2025-07-23 09:37:01 -03:00
Salvador E. Tropea f0cf78326c [Requirements][Added] SeCoNoHe 2025-07-23 09:31:14 -03:00
Salvador E. Tropea d387703101 Migrated to use SeCoNoHe 2025-07-23 09:20:16 -03:00
Salvador E. Tropea 1d4782516c Bumped version to 1.1.1 2025-07-21 09:39:57 -03:00
Salvador E. Tropea ad8ad49d9c [Examples][Added] Quick versions
They download the audio from internet
In most cases just uses 10 seconds to make it faster
2025-07-20 19:17:10 -03:00
Salvador E. Tropea 0b425c9bd1 [Demucs][Removed] Normalization
Not really needed
2025-07-18 11:03:54 -03:00
Salvador E. Tropea a377c72fac [Examples][Added] Demix and Remix example
Also added links to the examples from the README
2025-07-18 10:58:29 -03:00
Salvador E. Tropea ab3b9d1c3d [Demucs][Added] Torchaudio model support
The one you get using HDEMUCS_HIGH_MUSDB_PLUS
Is an HDemucs model, the code in TorchAudio looks old, perhaps
simplified with incompatible renames.
2025-07-12 13:49:32 -03:00
Salvador E. Tropea c87923faeb Bumped version to 1.1.0 2025-07-11 18:04:42 -03:00
Salvador E. Tropea 64070adf77 [Tool] Made all executable 2025-07-11 17:52:20 -03:00
Salvador E. Tropea b8ee1856c5 [Demix][Fixed] Handling of not generated stems 2025-07-11 17:51:45 -03:00
Salvador E. Tropea f973796f07 [Demucs][Replaced] Julius up/down sampler by torchaudio
Has less distortion.
2025-07-11 17:14:05 -03:00
Salvador E. Tropea 9a1297beaf [Demucs][Logger][Added] Signature
And made more robust the sync between sig and kwargs
2025-07-11 17:12:42 -03:00
Salvador E. Tropea 93856c88eb [Demucs][Logger][Fixed] Wiener iterations default 2025-07-11 17:12:02 -03:00
Salvador E. Tropea 75ed37ed3b [Fixed] Restored wiener
Used by the mdx.safetensors file, one HDemucs uses CaC and the
other Wiener
2025-07-11 17:09:35 -03:00
Salvador E. Tropea 08fb457a9d [Demucs][Removed] Weiner filtering support
Currently unused, will be restored only if we find a model using it
2025-07-11 10:01:33 -03:00
Salvador E. Tropea dfcce31318 [Demucs][Logger][Added] Better STFT framework log 2025-07-11 09:59:04 -03:00
Salvador E. Tropea 499aac5e80 [Demucs] Better logger
- Better code
- Better output
2025-07-11 09:24:53 -03:00
Salvador E. Tropea 6760256cd7 [Added] Debug information about Demucs
So we can know the architecture used
2025-07-10 13:08:18 -03:00
79 changed files with 912 additions and 1118 deletions
@@ -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:
+1
View File
@@ -14,3 +14,4 @@ __no__
models/.catalog.csv
models/*.yaml
0LEEME
sync.sh
+15
View File
@@ -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.
+43 -15
View File
@@ -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
View File
@@ -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

+28
View File
@@ -0,0 +1,28 @@
# Audio Separation assets
## Icon
![Icon 1024x1024](AudioSeparation_1024.jpg)
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:
![Icon 400x400](AudioSeparation_400.jpg)
## Banner
![Banner 21:9](audioseparation_logo_21_9.jpg)
Image composition
Binary file not shown.

After

Width:  |  Height:  |  Size: 69 KiB

File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
audioseparation_logo.jpg
File diff suppressed because one or more lines are too long
+1
View File
@@ -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
View File
@@ -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
View File
@@ -0,0 +1 @@
audioseparation_logo.jpg
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
audioseparation_logo.jpg
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
audioseparation_logo.jpg
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
audioseparation_logo.jpg
File diff suppressed because one or more lines are too long
-57
View File
@@ -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;
}
});
},
});
-32
View File
@@ -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
});
});
},
});
+30
View File
@@ -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
View File
@@ -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"
+1
View File
@@ -3,6 +3,7 @@ torchaudio
numpy
safetensors
tqdm
seconohe>=1.0.2
# Optionals:
# requests
# colorama
-220
View File
@@ -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)
View File
-113
View File
@@ -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()
-44
View File
@@ -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}")
-46
View File
@@ -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}")
-206
View File
@@ -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
-104
View File
@@ -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
-128
View File
@@ -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()
+13
View File
@@ -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]
+322
View File
@@ -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
+9 -8
View File
@@ -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
View File
@@ -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)
+80
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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: