[Added] Secondary tools
This commit is contained in:
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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
@@ -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)
|
||||
@@ -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)
|
||||
Executable
+226
@@ -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()
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user