[Added] mypy hook

Using local mypy to avoid pulling "yet another huge PyTorch copy"
This commit is contained in:
Salvador E. Tropea
2025-07-16 11:26:27 -03:00
parent 42c8e21444
commit 111b58a615
6 changed files with 65 additions and 11 deletions
+23
View File
@@ -41,3 +41,26 @@ repos:
# "--check-hidden"
]
# You can create a .codespellignore file with one word per line for words to ignore.
# --- Use a "local" hook for mypy ---
- repo: local
hooks:
- id: mypy
name: mypy
# The command to execute. It's expected to be found in the PATH
# of the environment where you run `git commit`.
entry: mypy
args: ["--explicit-package-bases", "source/"]
# 'system' tells pre-commit to find the command in the current environment
# instead of building an isolated one.
language: system
# Specify which files to run on.
types: [python]
pass_filenames: false
# You might still need args if not configured elsewhere, but a mypy.ini is better.
# args: [--config-file=mypy.ini]
# `pass_filenames: false` can be useful here if you want mypy to analyze the
# whole project as configured in mypy.ini, rather than just the changed files.
# Try without it first.
# pass_filenames: false
+22
View File
@@ -0,0 +1,22 @@
[mypy]
# Basic configuration
python_version = 3.11
warn_return_any = True
warn_unused_ignores = True
mypy_path = source/
# This is often needed when starting out, especially with libraries that lack stubs
ignore_missing_imports = True
# You can get stricter over time by removing ignore_missing_imports
# and setting flags like these:
# disallow_untyped_defs = True
# disallow_any_unimported = True
# Tell mypy to ignore errors from third-party libraries if they are not typed
[mypy-torch.*]
ignore_missing_imports = True
[mypy-torchaudio.*]
ignore_missing_imports = True
[mypy-pytest.*]
ignore_missing_imports = True
+3 -1
View File
@@ -7,6 +7,7 @@
import torch
import torchaudio
import torchaudio.transforms as T
from typing import Optional, Dict, Any
from .utils.aligner import AudioBatchAligner
from .utils.logger import main_logger
from .utils.misc import parse_time_to_seconds, parse_note_to_frequency
@@ -429,7 +430,7 @@ class AudioBlend:
UNIQUE_NAME = "SET_AudioBlend"
DISPLAY_NAME = "Audio Blend"
def blend_audio(self, audio1: dict, gain1: float, gain2: float, audio2: dict = None):
def blend_audio(self, audio1: dict, gain1: float, gain2: float, audio2: Optional[Dict[str, Any]] = None):
if audio2 is None:
# Handle the simple case where audio2 is not provided
logger.info(f"Blending audio1 only with gain {gain1}.")
@@ -609,6 +610,7 @@ class AudioTestSignalGenerator:
waveform_multichannel = torch.zeros((batch_size, channels, num_samples))
# Final clip to ensure [-1, 1] range, as some operations might exceed it slightly
assert waveform_multichannel is not None, "Waveform should have been created by now"
final_waveform = torch.clamp(waveform_multichannel, -1.0, 1.0)
logger.info(f"Generated '{waveform_type}' signal. Shape: {final_waveform.shape}, SR: {sample_rate}Hz")
+15 -8
View File
@@ -2,21 +2,28 @@
# Copyright (c) 2025 Instituto Nacional de Tecnologïa Industrial
# License: GPL-3.0
# Project: ComfyUI-AudioBatch
from __future__ import annotations # Good practice
import logging
import os
import sys
import logging
from typing import Any
from .misc import NODES_NAME, NODES_DEBUG_VAR
no_colorama = False
# 1. Initialize variables with the `Any` type.
# This tells mypy not to make assumptions about their specific class.
Fore: Any
Back: Any
Style: Any
# 2. Perform the runtime import logic as before.
try:
from colorama import init as colorama_init, Fore, Back, Style
except ImportError:
no_colorama = True
# If colorama isn't installed use an ANSI basic replacement
if no_colorama:
from .ansi import Fore, Back, Style # noqa: F811
else:
colorama_init()
except ImportError:
# If colorama is not available, import our fallback.
# mypy will now allow this assignment because the variables were declared as Any.
from .ansi import Fore, Back, Style
class CustomFormatter(logging.Formatter):
+1 -1
View File
@@ -144,7 +144,7 @@ def get_peak_frequency(waveform: torch.Tensor, sample_rate: int) -> float:
peak_index = torch.argmax(torch.abs(fft_result))
# Get the frequency corresponding to that peak index
peak_freq = freq_bins[peak_index].item()
return peak_freq
return float(peak_freq)
def test_resampler_preserves_frequency_content(resampler_node):
@@ -42,7 +42,7 @@ def get_peak_frequency(waveform: torch.Tensor, sample_rate: int) -> float:
peak_index = torch.argmax(torch.abs(fft_result))
# Get the frequency corresponding to that peak index
peak_freq = freq_bins[peak_index].item()
return peak_freq
return float(peak_freq)
# --- Test Cases ---