[Added] Secondary tools

This commit is contained in:
Salvador E. Tropea
2025-07-02 08:35:47 -03:00
parent e0b2bde974
commit d9fa1bf49f
8 changed files with 1296 additions and 0 deletions
+205
View File
@@ -0,0 +1,205 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Tool to convert all ONNX MDX-Net files into a safetensors
# Run it using: python tool/batch_convert.py
#
# Created using Gemini 2.5 Pro
import argparse
import json
import os
import re
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
def parse_converter_output(output):
"""Parses the verbose output of onnx2safetensors.py to extract parameters."""
params = {}
try:
params['dim_f'] = int(re.search(r"Frequency Dimension \(dim_f\):\s*(\d+)", output).group(1))
params['channels'] = int(re.search(r"Base Channels \(ch\):\s*(\d+)", output).group(1))
params['stages'] = int(re.search(r"U-Net Stages:\s*(\d+)", output).group(1))
params['params'] = int(re.search(r"Total parameters:\s*(\d+)", output).group(1))
except (AttributeError, TypeError):
# This happens if the output did not contain the expected lines
main_logger.error("No detection from convert.")
sys.exit(3)
return params
def main(args):
logger_set_standalone(args)
# 1. Load the JSON metadata file
model_db = load_known_models(args.json_file)
if model_db is None:
main_logger.error("No model database available.")
model_db = {}
# Ensure source and destination directories exist
if not os.path.isdir(args.source_dir):
main_logger.error(f"Source directory not found at '{args.source_dir}'.")
sys.exit(3)
os.makedirs(args.dest_dir, exist_ok=True)
# 2. Get list of files to convert
onnx_files = sorted([f for f in os.listdir(args.source_dir) if f.endswith('.onnx')])
if args.test:
main_logger.info("\n--- RUNNING IN TEST MODE ---")
onnx_files = ["Kim_Vocal_2.onnx", "kuielab_b_drums.onnx"]
for fn in onnx_files:
if not os.path.exists(os.path.join(args.source_dir, fn)):
main_logger.error(f"Test file '{fn}' not found in source directory.")
sys.exit(3)
main_logger.info(f"\nFound {len(onnx_files)} ONNX files to process.")
# 3. Process each file
for filename in onnx_files:
main_logger.info(f"\n{'='*50}")
main_logger.info(f"Processing '{filename}'...")
source_path = os.path.join(args.source_dir, filename)
# Check if destination file already exists
dest_filename = os.path.splitext(filename)[0] + ".safetensors"
dest_path = os.path.join(args.dest_dir, dest_filename)
if os.path.exists(dest_path):
main_logger.info(f"Destination file '{dest_path}' already exists. Skipping.")
continue
# Get file hash
file_hash = get_hash(source_path)
# Verify hash exists in JSON database
if file_hash not in model_db:
main_logger.error(f"Hash '{file_hash}' for file '{filename}' not found in JSON database.")
main_logger.error("Please add the model's metadata to the JSON file before converting.")
sys.exit(3)
model_info = model_db[file_hash]
main_logger.info(f"Found metadata for '{model_info.get('name', 'N/A')}'. Hash: {file_hash}")
# Prepare the command to call the converter tool
command = [
sys.executable,
args.converter_script,
source_path,
"-o", dest_path,
"-m", args.model_location,
"-j", json.dumps(model_info),
"-v" # Always use verbose mode to capture the parameters
]
main_logger.debug(f"Executing command: {' '.join(command)}")
# 4. Call the converter and validate its output
try:
result = subprocess.run(command, capture_output=True, text=True, check=True, encoding='utf-8')
main_logger.debug("--- Converter Output ---")
main_logger.debug(result.stdout)
if result.stderr:
main_logger.error("--- Converter Errors ---")
main_logger.error(result.stderr)
# 5. Parse output and compare parameters
detected_params = parse_converter_output(result.stdout)
# Get expected params from JSON, using defaults for missing values
expected_params = {
'dim_f': model_info.get("mdx_dim_f_set", 3072),
'channels': model_info.get("channels", 48),
'stages': model_info.get("stages", 5),
'params': model_info.get("params", 0)
}
main_logger.info(f"Comparing parameters: Expected {expected_params} vs Detected {detected_params}")
if detected_params == expected_params:
main_logger.info(f"✅ SUCCESS: Parameters match. Conversion for '{filename}' is valid.")
# 6. Compute the hash for the new safetensors file
main_logger.info(f"Computing hash for new file '{dest_path}'...")
new_hash = get_hash(dest_path)
main_logger.info(f"New file hash: {new_hash}")
# Find and remove any old entry that points to the same safetensors filename
old_hash_to_remove = None
for hash_key, data in model_db.items():
if data.get("name") == dest_filename:
old_hash_to_remove = hash_key
break
if old_hash_to_remove:
main_logger.info(f"Removing old database entry for '{dest_filename}' (hash: {old_hash_to_remove}).")
del model_db[old_hash_to_remove]
# Add the new entry, copying relevant metadata from the ONNX entry
main_logger.info(f"Adding new entry to database for hash: {new_hash}")
# 1. Start with a deep copy of the original ONNX model's metadata
new_entry = model_info.copy()
# 2. Update the keys with our new, verified information
new_entry["name"] = dest_filename # Update the filename
new_entry["mdx_dim_f_set"] = detected_params['dim_f'] # Update with detected value
new_entry["channels"] = detected_params['channels'] # Update with detected value
new_entry["stages"] = detected_params['stages'] # Update with detected value
new_entry["params"] = detected_params['params'] # Update with detected value
# Adjust some values to match the conversion
new_entry["file_t"] = "safetensors"
new_entry["download"] = "Main/MDX"
# 3. Add the new hash to the database
model_db[new_hash] = new_entry
else:
main_logger.error(f"Parameter mismatch for '{filename}'.")
main_logger.error("Deleting incorrect output file.")
if os.path.exists(dest_path):
os.remove(dest_path)
sys.exit(3)
except subprocess.CalledProcessError as e:
main_logger.error(f"onnx2safetensors.py script failed for '{filename}'.")
main_logger.error("--- STDOUT ---")
print(e.stdout)
main_logger.error("--- STDERR ---")
print(e.stderr)
main_logger.error("----------------")
if not args.test or args.always_save_db: # Only save the database if not in test mode
main_logger.info("\n--- Saving updated JSON database ---")
try:
save_known_models(model_db, args.json_file)
except Exception:
sys.exit(3)
else:
main_logger.info("\n--- Test mode finished. JSON database was NOT modified. ---")
if __name__ == "__main__":
parser = argparse.ArgumentParser(
description="A batch processing tool to convert ONNX models to Safetensors using a metadata file.",
formatter_class=argparse.ArgumentDefaultsHelpFormatter
)
parser.add_argument('--source_dir', type=str, default='models', help="Directory containing the input .onnx files.")
parser.add_argument('--dest_dir', type=str, default='models/new', help="Directory to save the output .safetensors files.")
parser.add_argument('--json_file', type=str, default=get_db_filename(), help="Path to the metadata JSON file.")
parser.add_argument('--test', action='store_true', help="Run in test mode, converting only 'Kim_Vocal_2.onnx'.")
parser.add_argument('--always_save_db', action='store_true', help="Save updated database even while in test mode.")
# Paths to the other tools in the same directory
parser.add_argument('--converter_script', type=str, default='tool/onnx2safetensors.py',
help="Source to convert the files.")
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)
args = parser.parse_args()
main(args)
+68
View File
@@ -0,0 +1,68 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Tool to download a model from the models Data Base.
# Run it using: python tool/download_model.py HASH
import argparse
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, get_download_url, cli_add_models_and_db
from source.utils.downloader import download_model
from source.utils.logger import main_logger, logger_set_standalone
from source.utils.misc import cli_add_verbose
def main(args):
logger_set_standalone(args)
# Load the JSON metadata file
model_db = load_known_models(args.json_file)
if model_db is None:
main_logger.error("No model database available.")
sys.exit(3)
# Is a valid hash?
if args.hash not in model_db:
main_logger.error(f"Nothing known about `{args.hash}`")
sys.exit(4)
d = model_db[args.hash]
# Check what we have
downloaded = hash_dir(args.models_dir)
if args.hash in downloaded:
main_logger.error(f"`{args.hash}` already downloaded as `{downloaded[args.hash]}`")
sys.exit(5)
# Check we can download it
url = get_download_url(d)
if url is None:
main_logger.error(f"`{args.hash}` can't be downloaded")
sys.exit(6)
# Download the file
name = d['name']
try:
download_model(url, args.models_dir, name, force_urllib=False)
except Exception as e:
main_logger.error(f"Failed to download {name} from {url}\n{e}")
raise
if __name__ == "__main__":
parser = argparse.ArgumentParser(
description="Downloads a model from the database.",
formatter_class=argparse.ArgumentDefaultsHelpFormatter
)
# --- File/Path Arguments ---
parser.add_argument('hash', type=str, help="Hash for the file to download.")
cli_add_models_and_db(parser)
# --- Control Arguments ---
cli_add_verbose(parser)
args = parser.parse_args()
main(args)
+239
View File
@@ -0,0 +1,239 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Tool to convert an ONNX MDX-Net file into a safetensors file
# Run it using: python tool/onnx2safetensors.py ONNX_FILE
#
# Created using Gemini 2.5 Pro, but with a lot of iterations and adjusts
import argparse
import json
import numpy as np
import onnx
import os
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
class OnnxGraph:
"""A helper class to hold ONNX graph information."""
def __init__(self, onnx_path: str):
main_logger.info("Loading ONNX graph...")
onnx_model = onnx.load(onnx_path)
self.nodes = list(onnx_model.graph.node)
self.weights = {t.name: torch.from_numpy(np.copy(onnx.numpy_helper.to_array(t))) for t in onnx_model.graph.initializer}
self.weighted_nodes = [n for n in self.nodes if any(inp in self.weights for inp in n.input)]
self.num_weighted_nodes = len(self.weighted_nodes)
main_logger.info(f"ONNX graph loaded. Found {self.num_weighted_nodes} nodes with weights.")
class AutoMapper:
"""
Performs the 'Dual Walk' using a reliable recursive traversal of the
PyTorch model's module hierarchy, eliminating the need for a forward-pass trace.
"""
def __init__(self, pytorch_model: nn.Module, onnx_graph: OnnxGraph, verbose: bool = False):
self.pytorch_model = pytorch_model
self.onnx_graph = onnx_graph
self.verbose = verbose
# Get an ordered list of PyTorch modules that have weights by traversing the module hierarchy.
# This is guaranteed to be in the same order as the state_dict.
self.pytorch_modules_with_params = [
m for m in self.pytorch_model.modules()
if len(list(m.children())) == 0 and len(list(m.parameters(recurse=False))) > 0
]
def add_target(self, map_list, key_iter, source, transpose=False):
target = next(key_iter)
map_list.append({'target': target, 'source': source, 'transpose': transpose})
if self.verbose > 1:
main_logger.debug(f" {source} -> {target}'")
def generate_map(self):
main_logger.info("Generating automatic weight map by walking module hierarchy...")
if len(self.pytorch_modules_with_params) != self.onnx_graph.num_weighted_nodes:
main_logger.error("🚨 Mismatch in number of weighted layers!")
main_logger.error(f"PyTorch model defines {len(self.pytorch_modules_with_params)} layers with weights.")
main_logger.error(f"ONNX file contains {self.onnx_graph.num_weighted_nodes} nodes with weights.")
main_logger.error("This likely means the model's parameters (dim_f, ch, num_stages) do not match the ONNX file.")
sys.exit(1)
map_list = []
# We can now reliably iterate through the state_dict keys.
key_iter = iter([k for k in self.pytorch_model.state_dict().keys() if 'num_batches_tracked' not in k])
for pt_module, onnx_node in zip(self.pytorch_modules_with_params, self.onnx_graph.weighted_nodes):
if self.verbose:
main_logger.debug(f" Mapping ONNX '{onnx_node.op_type}' -> PyTorch '{pt_module.__class__.__name__}'")
onnx_param_names = [name for name in onnx_node.input if name in self.onnx_graph.weights]
if isinstance(pt_module, nn.Linear):
self.add_target(map_list, key_iter, onnx_param_names[0], True)
elif isinstance(pt_module, (nn.Conv2d, nn.ConvTranspose2d)):
self.add_target(map_list, key_iter, onnx_param_names[0])
if pt_module.bias is not None:
self.add_target(map_list, key_iter, onnx_param_names[1])
elif isinstance(pt_module, nn.BatchNorm2d):
# The order in the state_dict is weight, bias, running_mean, running_var
# The order in the ONNX node is scale, B, mean, var
# They correspond 1-to-1
for i in range(4):
self.add_target(map_list, key_iter, onnx_param_names[i])
main_logger.info("Automatic map generated successfully.")
return map_list
def verify_parameter(params, name, value, description):
""" Check if a parameter is already known.
If known check our detection is ok.
Otherwise add it """
in_params = params.get(name)
if in_params:
if in_params != value:
main_logger.error(f"{description} mismatch ({in_params} vs {value})")
sys.exit(3)
else:
params[name] = value
def convert(args):
"""Main conversion function driven by command line arguments."""
main_logger.info(f"Starting conversion for '{args.input_file}'...")
onnx_graph = OnnxGraph(args.input_file)
# Auto-detect parameters
first_matmul_weight_name = next((n.input[1] for n in onnx_graph.nodes if n.op_type == 'MatMul'), None)
if not first_matmul_weight_name:
raise ValueError("Could not find a MatMul node to detect dim_f.")
dim_f = onnx_graph.weights[first_matmul_weight_name].shape[0]
first_conv_weight_name = next((n.input[1] for n in onnx_graph.nodes if n.op_type == 'Conv'), None)
if not first_conv_weight_name:
raise ValueError("Could not find a Conv node to detect channels.")
ch = onnx_graph.weights[first_conv_weight_name].shape[0]
num_stages = 5 if onnx_graph.num_weighted_nodes > 80 else 4
total_weight_params = 0
for tensor in onnx_graph.weights.values():
total_weight_params += tensor.numel() # numel() gives the total number of elements
main_logger.info("\n--- Detected Model Parameters ---")
main_logger.info(f" Frequency Dimension (dim_f): {dim_f}")
main_logger.info(f" Base Channels (ch): {ch}")
main_logger.info(f" U-Net Stages: {num_stages}")
main_logger.info(f" Total parameters: {total_weight_params}")
TargetModelClass = import_model_class(args.model_location)
target_model = TargetModelClass(dim_f=dim_f, ch=ch, num_stages=num_stages)
auto_mapper = AutoMapper(target_model, onnx_graph, args.verbose)
mapping = auto_mapper.generate_map()
main_logger.info("\n--- Starting Final Weight Transfer ---")
new_state_dict = target_model.state_dict()
transfer_count = 0
for item in mapping:
target_key, source_key, needs_transpose = item['target'], item['source'], item['transpose']
if target_key in new_state_dict:
if source_key in onnx_graph.weights:
source_tensor = onnx_graph.weights[source_key]
if needs_transpose:
source_tensor = source_tensor.t()
if new_state_dict[target_key].shape == source_tensor.shape:
new_state_dict[target_key].data.copy_(source_tensor)
transfer_count += 1
else:
main_logger.warning(f" [!] SHAPE MISMATCH for '{target_key}': Target {new_state_dict[target_key].shape} "
f"vs Source {source_tensor.shape}")
else:
main_logger.error(f" [!] ERROR: Raw weight key '{source_key}' not found.")
else:
main_logger.warning(f" [!] Warning: Target key '{target_key}' not found.")
# Finalize and Save
state_dict_to_save = target_model.state_dict()
final_state_dict = {
key: tensor for key, tensor in state_dict_to_save.items()
if 'num_batches_tracked' not in key
}
total_params = len(new_state_dict)
num_non_param_buffers = total_params - len(final_state_dict)
loadable_params = len(final_state_dict)
main_logger.info(f"\nSuccessfully transferred {transfer_count} / {loadable_params} tensors.")
if transfer_count < loadable_params:
main_logger.warning("Some weights were not transferred. Check for errors above.")
else:
main_logger.info(" All loadable parameters and buffers were successfully transferred.")
main_logger.info(f" ({num_non_param_buffers} non-persistent buffers like 'num_batches_tracked' "
"will be excluded from the final file.)")
output_path = args.output_file or os.path.splitext(args.input_file)[0] + ".safetensors"
# Add some metadata to the file
hyperparameters = args.metadata if args.metadata is not None else {}
# Verify dim_f
verify_parameter(hyperparameters, "mdx_dim_f_set", dim_f, "Frequency Dimension (dim_f)")
# Verify channels
verify_parameter(hyperparameters, "channels", ch, "Base Channels (ch)")
# Verify stages
verify_parameter(hyperparameters, "stages", num_stages, "U-Net Stages")
# Verify total parameters
verify_parameter(hyperparameters, "params", total_weight_params, "Total parameters")
# Fix entries to match the conversion
hyperparameters["name"] = os.path.basename(output_path)
hyperparameters["download"] = "Main/MDX"
hyperparameters["download"] = get_download_url(hyperparameters) # Convert to something usable outside our project
hyperparameters["project"] = "https://github.com/set-soft/AudioSeparation"
hyperparameters["file_t"] = "safetensors"
# They must be strings
model_hyperparameters = {k: str(v) for k, v in hyperparameters.items()}
main_logger.info(f"\nSaving model weights to '{output_path}'...")
os.makedirs(os.path.dirname(output_path), exist_ok=True)
# Save the filtered state_dict
save_file(tensors=final_state_dict, filename=output_path, metadata=model_hyperparameters)
main_logger.info("\n✅ Conversion complete!")
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Smart tool to convert specific MDX-Net ONNX models to PyTorch Safetensors.",
formatter_class=argparse.RawTextHelpFormatter)
parser.add_argument('input_file', type=str, help="Path to the input ONNX model file.")
parser.add_argument('-o', '--output_file', type=str, default=None,
help="Path for the output .safetensors file.\n(default: same as input with .safetensors extension)")
parser.add_argument('-m', '--model_location', type=str, default='source/inference/MDX_Net.py:MDX_Net',
help="Location of the PyTorch model class, in 'filename:ClassName' format.\n"
"(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)
args = parser.parse_args()
logger_set_standalone(args)
if args.metadata is not None:
try:
args.metadata = json.loads(args.metadata)
except Exception as e:
main_logger.error(f"Failed to parse the metadata: {e}")
sys.exit(2)
if not isinstance(args.metadata, dict):
main_logger.error(f"Metadata must be a dict, not {type(args.metadata)}")
sys.exit(2)
convert(args)
+146
View File
@@ -0,0 +1,146 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Tool to show PyTorch class i.e:
# python tool/show_class.py -m source/inference/MDX_Net.py:MDX_Net
import argparse
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
def load_class(args):
"""
Loads the structure of a specified PyTorch model class
"""
main_logger.info("--- Loading the class ---\n")
main_logger.info(f"Model Class: {args.model_location}")
main_logger.info(f"Parameters: dim_f={args.dim_f}, ch={args.ch}, num_stages={args.num_stages}")
# 1. Dynamically import the specified model class
TargetModelClass = import_model_class(args.model_location)
# 2. Instantiate the model with the provided parameters
try:
pytorch_model = TargetModelClass(dim_f=args.dim_f, ch=args.ch, num_stages=args.num_stages)
except Exception as e:
main_logger.error("Could not instantiate the model class with the given parameters.")
main_logger.error(f"Please check if the class '{args.model_location.split(':')[1]}' "
"accepts 'dim_f', 'ch', and 'num_stages'.")
main_logger.error(f"Original error: {e}")
sys.exit(1)
return pytorch_model
def show(model):
"""
Prints the structure of a specified PyTorch model class
"""
# Print the PyTorch model structure to the log
# The f-string ensures the multi-line output of the model is captured in the log
main_logger.info(f"\n--- PyTorch Model Structure ---\n\n{model}")
def export_onnx(model, dim_f, file):
import torch
main_logger.info("\n--- Exporting to ONNX ---")
# 1. Create a dummy input with the correct shape
dummy_input = torch.randn(1, 4, dim_f, 256)
model.eval()
# 2. Export the model
try:
torch.onnx.export(
model,
dummy_input,
file,
export_params=False, # We only care about the graph structure, not weights
opset_version=11,
input_names=['input'],
output_names=['output'],
dynamic_axes={'input': {0: 'batch_size'}, 'output': {0: 'batch_size'}}
)
main_logger.info(f"\n✅ SUCCESS! Model successfully exported to '{file}'")
except Exception as e:
main_logger.error(f"\n❌ FAILURE during export: {e}")
def show_keys(model):
main_logger.info("\n--- PyTorch Model Keys ---\n\n")
for key in model.state_dict().keys():
main_logger.info(key)
def show_compact(module, parent_name='model', indent=0, index=0):
"""
Recursively walks the PyTorch model and prints its structure
in a format similar to our ONNX analysis.
"""
# Get the module's class name and a formatted name with its parent
module_name = module.__class__.__name__
full_name = f"{parent_name}.{module_name}" if parent_name else module_name
# Print containers like ModuleList or Sequential differently
if len(list(module.children())) > 0 and not isinstance(module, nn.Sequential):
main_logger.info(f"[---] {' ' * indent}{full_name}: {module_name}")
# Iterate through named children to maintain order
for name, child_module in module.named_children():
child_full_name = f"{parent_name}.{name}"
# If the child is a leaf node (like Conv2d, Linear, etc.)
if len(list(child_module.children())) == 0:
main_logger.info(f"[{index:03d}] {' ' * (indent+1)}{child_full_name}: {child_module.__class__.__name__}")
index += 1
else:
# If the child is another container, recurse
index = show_compact(child_module, child_full_name, indent + 1, index)
return index
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Shows the current class layers and/or keys",
formatter_class=argparse.RawTextHelpFormatter)
parser.add_argument('-m', '--model_location', type=str, default='source/inference/MDX_Net.py:MDX_Net',
help="Location of the PyTorch model class, in 'filename:ClassName' format.\n"
"(default: source/inference/MDX_Net.py:MDX_Net)")
parser.add_argument('-d', '--dim_f', type=int, default=3072,
help="The frequency dimension (dim_f) of the model. (default: 3072)")
parser.add_argument('-c', '--ch', type=int, default=48,
help="The base channel count (ch) of the model. (default: 48)")
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.")
args = parser.parse_args()
logger_set_standalone(args)
model = load_class(args)
if args.no_show:
show(model)
# Optional ONNX export
if args.export_onnx:
export_onnx(model, args.dim_f, args.export_onnx)
# Optional print keys
if args.keys:
show_keys(model)
# Optional compact form
if args.compact:
main_logger.info("\n--- PyTorch Compact Model Structure ---\n")
show_compact(model)
+133
View File
@@ -0,0 +1,133 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Tool to show the models Data Base.
# Is used to do adjusts to the data base.
# Run it using: python tool/show_db.py
import argparse
import pprint
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
# Do nothing, you can apply some change here
def apply_process(model_db):
return model_db, False
# Example of processing
# PARAMS = {"3072/48/5": 16684228,
# "2560/48/5": 14763012,
# "2048/48/5": 13191108,
# "2048/32/5": 7420548,
# "2048/32/4": 5478276}
#
# def apply_process(model_db):
# for k, v in model_db.items():
# if 'name' not in v:
# continue
# main_logger.info(f"Processing {v['desc']}")
# key = f"{v['mdx_dim_f_set']}/{v['channels']}/{v['stages']}"
# try:
# v['params'] = PARAMS[key]
# except KeyError:
# print(f"No {key}")
# raise
# return model_db, True
# Example of processing
# def apply_process(model_db):
# modified = False
# for k, v in model_db.items():
# if 'name' not in v:
# continue
# main_logger.info(f"Processing {v['desc']}")
# if 'channels' not in v:
# v['channels'] = 48
# main_logger.info("- Explicit 48 channels")
# if 'stages' not in v:
# v['stages'] = 5
# main_logger.info("- Explicit 5 stages")
# v['model_t'] = 'MDX'
# main_logger.info("- Explicit model_t MDX")
# v['file_t'] = os.path.splitext(v['name'])[1][1:]
# main_logger.info(f"- Explicit file_t {v['file_t']}")
# modified = True
# return model_db, modified
def main(args):
logger_set_standalone(args)
# 1. Load the JSON metadata file
model_db = load_known_models(args.json_file)
if model_db is None:
main_logger.error("No model database available.")
sys.exit(3)
main_logger.info("\n--- Current DB ---\n")
pprint.pprint(model_db)
# 2. Apply any hardcoded processing changes (if any)
model_db, modified = apply_process(model_db)
# 3. Apply command-line filters if any were provided
filters_active = args.primary_stem or args.model_t or args.file_t
if filters_active:
main_logger.info("\n--- Hashing downloaded models ---\n")
hashes = hash_dir(args.models_dir)
pprint.pprint(hashes)
main_logger.info("\n--- Filtered DB ---\n")
# The get_models function can handle lists/sets of values directly
models_data, filtered_results = get_models(
primary_stem=set(args.primary_stem) if args.primary_stem is not None else None,
model_t=set(args.model_t) if args.model_t is not None else None,
file_t=set(args.file_t) if args.file_t is not None else None,
json_path=args.json_file,
downloaded=hashes
)
pprint.pprint(filtered_results)
pprint.pprint(models_data)
# 4. If the apply_process function modified the DB, show it and save it
if modified:
main_logger.info("\n--- Modified DB (to be saved) ---\n")
pprint.pprint(model_db)
main_logger.info("\n--- Saving updated JSON database ---")
try:
save_known_models(model_db, args.json_file)
main_logger.info("Save successful.")
except Exception as e:
main_logger.error(f"Failed to save database: {e}")
sys.exit(3)
if __name__ == "__main__":
parser = argparse.ArgumentParser(
description="Show the content of the models database.",
formatter_class=argparse.ArgumentDefaultsHelpFormatter
)
# --- File/Path Arguments ---
cli_add_models_and_db(parser)
# --- Filtering Arguments ---
parser.add_argument('--primary_stem', action='append',
help="Filter by primary stem. Can be used multiple times (e.g., --primary_stem Vocals).")
parser.add_argument('--model_t', action='append',
help="Filter by model type. Can be used multiple times (e.g., --model_t MDX --model_t Demucs).")
parser.add_argument('--file_t', action='append',
help="Filter by file type. Can be used multiple times (e.g., --file_t safetensors).")
# --- Control Arguments ---
cli_add_verbose(parser)
args = parser.parse_args()
main(args)
+219
View File
@@ -0,0 +1,219 @@
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Tool to show the ONNX structure, i.e:
# python tool/show_onnx.py model.onnx
import argparse
import numpy as np
import onnx
from onnx import shape_inference # Import the shape inference module
import torch
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
def print_onnx_nodes_and_weights(onnx_model_path):
"""
Analyzes and prints a detailed map of an ONNX model, including
type and shape information for all tensors, populated via shape inference.
"""
original_model = onnx.load(onnx_model_path)
# --- 1. Run Shape Inference to populate value_info ---
# This is the crucial step that adds missing type/shape info.
main_logger.info("Running ONNX shape inference...")
inferred_model = shape_inference.infer_shapes(original_model)
graph = inferred_model.graph
main_logger.info("Shape inference complete.")
# --- 2. Create a comprehensive map of all tensors in the graph ---
value_info_all = {
value_info.name: value_info
for value_info in list(graph.input) + list(graph.value_info) + list(graph.output)
}
main_logger.info("\n--- Model Inputs ---")
for input_tensor in graph.input:
tensor_type = onnx.helper.printable_type(input_tensor.type)
main_logger.info(f"Name: {input_tensor.name}, Type: {tensor_type}")
main_logger.info("\n--- Model Outputs ---")
for output_tensor in graph.output:
tensor_type = onnx.helper.printable_type(output_tensor.type)
main_logger.info(f"Name: {output_tensor.name}, Type: {tensor_type}")
# --- 3. Prepare data for analysis ---
nodes = list(graph.node)
weights = {t.name: torch.from_numpy(np.copy(onnx.numpy_helper.to_array(t))) for t in graph.initializer}
node_outputs = {out: n.name for n in nodes for out in n.output}
main_logger.info(f"\n--- Map for the ONNX graph with {len(nodes)} nodes ---")
for ni, n in enumerate(nodes):
has_weight = " (W)" if any(inp in weights for inp in n.input) else ""
main_logger.info(f"[{ni:03d}] {n.name} [{n.op_type}]{has_weight}")
# Print detailed inputs with type info
main_logger.info(" Inputs:")
for ninp, inp_name in enumerate(n.input):
from_txt, type_txt = "", ""
if inp_name in weights:
from_txt = f"weight initializer [name: {inp_name}]"
type_txt = str(weights[inp_name].shape)
elif inp_name in node_outputs:
from_txt = f"output of node '{node_outputs[inp_name]}'"
elif inp_name in [i.name for i in graph.input]:
from_txt = "graph input"
else:
from_txt = "unknown source"
if inp_name in value_info_all:
type_txt = onnx.helper.printable_type(value_info_all[inp_name].type)
main_logger.info(f" - [{ninp}] '{inp_name}' (from {from_txt}) -> Type: {type_txt}")
# Print detailed outputs with type info
main_logger.info(" Outputs:")
for nout, out_name in enumerate(n.output):
type_txt = (onnx.helper.printable_type(value_info_all.get(out_name, "").type)
if out_name in value_info_all else "N/A")
main_logger.info(f" - [{nout}] '{out_name}' -> Type: {type_txt}")
final_output_name = graph.output[0].name
main_logger.info(f"\nThe final graph output '{final_output_name}' is from node "
f"'{node_outputs.get(final_output_name, 'N/A')}'\n")
# --- Create a dummy input tensor by directly reading the graph's structured data ---
primary_input_info = graph.input[0]
# 1. Get the shape
shape = []
# The shape is stored in the 'dim' attribute of the tensor type
for dimension in primary_input_info.type.tensor_type.shape.dim:
# Check if the dimension has a fixed integer value or is a dynamic parameter
if dimension.HasField('dim_value'):
shape.append(dimension.dim_value)
else:
# For dynamic dimensions (like 'batch_size'), use 1 as a placeholder
shape.append(1)
# 2. Get the data type
try:
# Get the integer enum for the element type (e.g., 1 for FLOAT)
elem_type_enum = primary_input_info.type.tensor_type.elem_type
# Use ONNX's official mapping to get the corresponding NumPy dtype
np_dtype = onnx.helper.tensor_dtype_to_np_dtype(elem_type_enum)
# Convert the NumPy dtype to a PyTorch dtype
torch_dtype = torch.from_numpy(np.array(0, dtype=np_dtype)).dtype
except (KeyError, AttributeError):
# Fallback to float32 if type is not specified or recognized
main_logger.warning("Could not determine input dtype from ONNX graph. Defaulting to float32.")
torch_dtype = torch.float32
if not shape:
main_logger.error("Could not determine input shape from ONNX graph.")
return None
main_logger.info(f"Generating a random tensor of shape {shape} and type {torch_dtype} for inference.")
# 3. Create the random tensor
torch.manual_seed(0)
dummy_input_tensor = torch.randn(shape, dtype=torch_dtype)
# --- 4. Calculate and Print Model Statistics ---
num_weighted_nodes = len([n for n in nodes if any(inp in weights for inp in n.input)])
num_weight_tensors = len(weights)
# Calculate the total number of parameters (like "13B")
total_params = 0
for tensor in weights.values():
total_params += tensor.numel() # numel() gives the total number of elements
# Format the total parameters into a human-readable string (e.g., 1.2M, 2.5B)
if total_params > 1_000_000_000:
params_str = f"{total_params / 1_000_000_000:.2f}B"
elif total_params > 1_000_000:
params_str = f"{total_params / 1_000_000:.2f}M"
elif total_params > 1_000:
params_str = f"{total_params / 1_000:.2f}K"
else:
params_str = f"{total_params}"
main_logger.info("\n--- Model Statistics ---\n")
main_logger.info(f"Number of layers with weights: {num_weighted_nodes}")
main_logger.info(f"Total number of weight/bias/parameter tensors: {num_weight_tensors}")
main_logger.info(f"Total model parameters: {total_params:,} (~{params_str})")
return dummy_input_tensor
def convert_model(file):
try:
from onnx2pytorch import ConvertModel
except Exception:
main_logger.error("Install onnx2pytorch")
sys.exit(3)
main_logger.info("--- Loading the ONNX ---\n")
# Load the ONNX model
onnx_model = onnx.load(file)
main_logger.info("--- Converting to PyTorch ---\n")
# Convert the ONNX model to a PyTorch model
pytorch_model = ConvertModel(onnx_model)
return pytorch_model
def show_converted(model):
main_logger.info(f"\n--- PyTorch Model Structure ---\n\n{model}")
def show_keys(model):
main_logger.info("\n--- PyTorch Model Keys ---\n\n")
for key in model.state_dict().keys():
main_logger.info(key)
def run_model(model, input):
main_logger.info("\n--- Running Forward Pass ---")
if input is None:
main_logger.error("Missing random input for inference run")
return
# Run the forward pass to trigger our interception prints
with torch.no_grad():
_ = model(input)
main_logger.info("\n--- Run Complete ---")
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Shows the ONNX layers",
formatter_class=argparse.RawTextHelpFormatter)
parser.add_argument('input_file', type=str, help="Path to the input ONNX model file.")
parser.add_argument('-c', '--show_converted', action='store_true', help="Show the class converted using onnx2pytorch.")
parser.add_argument('-k', '--keys', action='store_true', help="Print the keys for the converted state_dict.")
parser.add_argument('-r', '--run', action='store_true',
help="Run un inference using the converted model.\n"
"Incompatible with -S")
parser.add_argument('-S', '--no_show', action='store_false', help="Don't print the ONNX structure.")
cli_add_verbose(parser)
args = parser.parse_args()
logger_set_standalone(args)
if args.run and not args.no_show:
main_logger.error("-r can't be used when -S is specified")
sys.exit(1)
dummy_input_tensor = print_onnx_nodes_and_weights(args.input_file) if args.no_show else None
if args.show_converted or args.keys or args.run:
model = convert_model(args.input_file)
if args.show_converted:
show_converted(model)
if args.keys:
show_keys(model)
if args.run:
run_model(model, dummy_input_tensor)
+226
View File
@@ -0,0 +1,226 @@
#!/usr/bin/python3
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Tool to fix some common style errors in code generated by Gemini
# Created by Gemini itself ;-)
import argparse
import os
import shutil # For file backup
import logging
# Setup logging
logger = logging.getLogger("StyleFixer")
logger.setLevel(logging.INFO) # Default level
# Create console handler and set level to debug
ch = logging.StreamHandler()
ch.setLevel(logging.DEBUG) # Let handler decide what to show based on logger's effective level
# Create formatter
formatter = logging.Formatter('%(asctime)s - %(name)s - %(levelname)s - %(message)s')
# Add formatter to ch
ch.setFormatter(formatter)
# Add ch to logger
logger.addHandler(ch)
def expand_tabs(line, tab_size=4):
"""Expands tabs to spaces, respecting tab stops."""
out_line = []
current_col = 0
for char in line:
if char == '\t':
spaces_to_add = tab_size - (current_col % tab_size)
out_line.append(' ' * spaces_to_add)
current_col += spaces_to_add
else:
out_line.append(char)
current_col += 1
return "".join(out_line)
def fix_indentation(line, tab_size=4):
"""Rounds leading whitespace indentation to the nearest multiple of tab_size."""
stripped_line = line.lstrip(' ')
if not stripped_line or stripped_line == line: # No leading spaces or empty line
return line
leading_spaces_count = len(line) - len(stripped_line)
remainder = leading_spaces_count % tab_size
if remainder == 0:
new_indent_count = leading_spaces_count
elif remainder <= tab_size / 2: # Prioritize rounding down or to current if exact multiple
new_indent_count = leading_spaces_count - remainder
else: # Round up
new_indent_count = leading_spaces_count + (tab_size - remainder)
new_indent_count = max(0, new_indent_count) # Ensure non-negative
if new_indent_count != leading_spaces_count:
logger.debug(f"Adjusting indent: from {leading_spaces_count} to {new_indent_count} for line: {line.rstrip()!r}")
return ' ' * new_indent_count + stripped_line
def process_line(line_num, line_content_with_eol, expand_tabs_enabled=True):
original_line_for_debug = line_content_with_eol # Keep for debugging if needed
line = line_content_with_eol
# 1. Expand tabs (if enabled)
if expand_tabs_enabled:
line = expand_tabs(line)
if line != original_line_for_debug and not line.isspace(): # Log only if actual content changed
logger.debug(f"L{line_num}: Tabs expanded.")
# 2. W291: Remove trailing whitespace (done before other checks that might rely on EOL)
# Also handles W293 (blank line contains whitespace) implicitly if the line becomes empty.
processed_line = line.rstrip()
if len(processed_line) < len(line.rstrip('\r\n')): # Compare lengths without EOL
logger.debug(f"L{line_num}: Trailing whitespace removed.")
# If the line became empty after rstrip, it's a truly blank line
if not processed_line and line.strip() == "":
# Return an empty string for a blank line, newline char will be added later if original had it
if line.endswith(('\n', '\r\n', '\r')):
return '\n'
return ""
# 3. Fix indentation (E111) - only if line is not blank
if processed_line:
indented_line = fix_indentation(processed_line)
if indented_line != processed_line: # Log if indentation changed
# Already logged inside fix_indentation with more detail
pass
processed_line = indented_line
# 4. E261: At least two spaces before inline comment
in_string_char = None
temp_part = ""
comment_starts_at = -1
for idx, char_val in enumerate(processed_line):
if in_string_char:
temp_part += char_val
if char_val == in_string_char:
if len(temp_part) >= 2 and temp_part[-2] == '\\':
pass
else:
in_string_char = None
elif char_val in ("'", '"'):
temp_part += char_val
in_string_char = char_val
elif char_val == '#':
comment_starts_at = idx
break # Found the first non-string comment char
else:
temp_part += char_val
if comment_starts_at != -1:
code_part = processed_line[:comment_starts_at]
comment_part = processed_line[comment_starts_at:] # Includes '#'
rstripped_code_part = code_part.rstrip()
# Only add spaces if there's actual code before the comment
if rstripped_code_part:
if not code_part.endswith(' '): # Needs fixing (less than 2 spaces)
new_line_with_comment_spacing = rstripped_code_part + ' ' + comment_part
if new_line_with_comment_spacing != processed_line:
logger.debug(f"L{line_num}: Adjusted inline comment spacing.")
processed_line = new_line_with_comment_spacing
# If rstripped_code_part is empty, it's a comment at the start of the line (after indent), leave it.
# Add back the newline character that rstrip might have removed, if the original line had one
# or if it's not an empty line that was purely whitespace
if line_content_with_eol.endswith('\n'):
return processed_line + '\n'
elif line_content_with_eol.endswith('\r\n'):
return processed_line + '\r\n'
elif line_content_with_eol.endswith('\r'):
return processed_line + '\r'
else: # Original line did not end with EOL
return processed_line
def process_file(filepath, expand_tabs_enabled=True, dry_run=False, no_backup=False):
logger.info(f"Processing {filepath}...")
try:
# Read with universal newlines mode, then splitlines to preserve EOLs correctly for rejoining
with open(filepath, 'r', encoding='utf-8', newline='') as f:
# original_content = f.read() # Reading all at once
# original_lines = original_content.splitlines(keepends=True) # This is better
original_lines = f.readlines() # Simpler, usually works fine
except Exception as e:
logger.error(f"Error reading file {filepath}: {e}")
return False
fixed_lines = []
changes_made = False
for i, line_content_with_eol in enumerate(original_lines):
fixed_line_content = process_line(i + 1, line_content_with_eol, expand_tabs_enabled)
fixed_lines.append(fixed_line_content)
if fixed_line_content != line_content_with_eol:
changes_made = True
logger.debug(f"L{i+1}: Original: {line_content_with_eol.rstrip()!r}")
logger.debug(f"L{i+1}: Fixed : {fixed_line_content.rstrip()!r}")
if changes_made:
if dry_run:
logger.info(f"Would fix style issues in {filepath} (dry run)")
else:
if not no_backup:
backup_dir = os.path.dirname(filepath)
backup_filename = "." + os.path.basename(filepath) + "~"
backup_path = os.path.join(backup_dir, backup_filename)
try:
shutil.copy2(filepath, backup_path) # copy2 preserves metadata
logger.info(f"Backup of original file created at {backup_path}")
except Exception as e:
logger.error(f"Failed to create backup for {filepath}: {e}. File not modified.")
return False
try:
# Write back lines using the original EOLs if possible, or common EOL
with open(filepath, 'w', encoding='utf-8', newline='') as f:
f.writelines(fixed_lines)
logger.info(f"Fixed style issues in {filepath}")
except Exception as e:
logger.error(f"Error writing to file {filepath}: {e}")
return False
else:
logger.info(f"No style issues to fix in {filepath}")
return True
def main():
parser = argparse.ArgumentParser(description="Fix common Python code style issues (Flake8).")
parser.add_argument("files", metavar="FILE", type=str, nargs='+',
help="Python files to process")
parser.add_argument("--no-expand-tabs", action="store_false", dest="expand_tabs",
help="Disable expansion of tabs to spaces.")
parser.add_argument("--dry-run", action="store_true",
help="Show what would be changed without modifying files.")
parser.add_argument("--no-backup", action="store_true",
help="Do not create a backup of the original file.")
parser.add_argument("-v", "--verbose", action="store_true",
help="Enable verbose (debug) logging.")
parser.set_defaults(expand_tabs=True)
args = parser.parse_args()
if args.verbose:
logger.setLevel(logging.DEBUG)
else:
logger.setLevel(logging.INFO)
for filepath in args.files:
if not os.path.isfile(filepath):
logger.warning(f"File not found: {filepath}. Skipping.")
continue
process_file(filepath, args.expand_tabs, args.dry_run, args.no_backup)
if __name__ == "__main__":
main()
+60
View File
@@ -0,0 +1,60 @@
#!/usr/bin/env python3
# Copyright (c) 2025 Salvador E. Tropea
# Copyright (c) 2025 Instituto Nacional de Tecnología Industrial
# License: GPLv3
# Project: ComfyUI-AudioSeparation
#
# Tool to get the hash used by UVR for a model i.e:
# python tool/uvr_hash.py model.onnx
import argparse
import os
# Local imports
import bootstrap # noqa: F401
from source.utils.logger import main_logger, logger_set_standalone
from source.db.hash import get_hash
def main():
"""Main function to run the command-line tool."""
parser = argparse.ArgumentParser(
description="""Compute a special MD5 hash of a file.
This tool calculates the hash of the last ~10MB of a file.
If the file is smaller than that, it hashes the entire file.
This method is often used for quick identification of large model files.
""",
# Makes the description formatting look nicer in the help text
formatter_class=argparse.RawTextHelpFormatter
)
parser.add_argument(
'files',
metavar='FILE',
nargs='+', # Accepts one or more file arguments
help='Path to the file(s) to hash.'
)
args = parser.parse_args()
args.verbose = 0
logger_set_standalone(args)
# Process each file provided on the command line
for filepath in args.files:
if not os.path.exists(filepath):
main_logger.error(f"File not found at '{filepath}'")
continue # Skip to the next file
if not os.path.isfile(filepath):
main_logger.error(f"Path '{filepath}' is a directory, not a file.")
continue # Skip to the next file
try:
file_hash = get_hash(filepath)
# Print in a standard format, similar to md5sum
main_logger.info(f"{file_hash} {filepath}")
except Exception as e:
main_logger.error(f"while processing file '{filepath}': {e}")
if __name__ == "__main__":
main()