diff --git a/tool/batch_convert.py b/tool/batch_convert.py new file mode 100644 index 0000000..824e852 --- /dev/null +++ b/tool/batch_convert.py @@ -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) diff --git a/tool/download_model.py b/tool/download_model.py new file mode 100644 index 0000000..4ba3510 --- /dev/null +++ b/tool/download_model.py @@ -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) diff --git a/tool/onnx2safetensors.py b/tool/onnx2safetensors.py new file mode 100644 index 0000000..c1403f8 --- /dev/null +++ b/tool/onnx2safetensors.py @@ -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) diff --git a/tool/show_class.py b/tool/show_class.py new file mode 100644 index 0000000..deb62ed --- /dev/null +++ b/tool/show_class.py @@ -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) diff --git a/tool/show_db.py b/tool/show_db.py new file mode 100644 index 0000000..9ff86a0 --- /dev/null +++ b/tool/show_db.py @@ -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) diff --git a/tool/show_onnx.py b/tool/show_onnx.py new file mode 100644 index 0000000..9a882fd --- /dev/null +++ b/tool/show_onnx.py @@ -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) diff --git a/tool/style_fixer.py b/tool/style_fixer.py new file mode 100755 index 0000000..39457d6 --- /dev/null +++ b/tool/style_fixer.py @@ -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() diff --git a/tool/uvr_hash.py b/tool/uvr_hash.py new file mode 100644 index 0000000..5e8798e --- /dev/null +++ b/tool/uvr_hash.py @@ -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()