diff --git a/.github/workflows/packaging.yml b/.github/workflows/packaging.yml index bd5332a..3e0ecd0 100644 --- a/.github/workflows/packaging.yml +++ b/.github/workflows/packaging.yml @@ -66,6 +66,7 @@ jobs: --add-data "examples;examples/" \ --add-data "assets;assets/" \ --add-data "locales;locales/" \ + --add-data "static;static/" \ --add-data "README.md;." \ --add-data "LICENSE;." \ --additional-hooks-dir build/hooks \ diff --git a/.gitignore b/.gitignore index 7aba46c..56cd7bf 100644 --- a/.gitignore +++ b/.gitignore @@ -11,6 +11,8 @@ build/* !build/auto-py-to-exe.json !build/hooks/ !build/hatch_build.py +static/vendor/* +!static/vendor/.gitkeep *.pyc *PitchLoader Output*.ustx *output*.ustx diff --git a/README.en.md b/README.en.md index 109e3c9..bc257d5 100644 --- a/README.en.md +++ b/README.en.md @@ -17,9 +17,7 @@ The current version supports importing the following expression parameters: * `Pitch Deviation (curve)` * `Tension (curve)` -

- -

+https://github.com/user-attachments/assets/4b5b7c15-947a-4f54-b80e-a14a9eefc86b > - *OpenUtau version used from [keirokeer/OpenUtau-DiffSinger-Lunai](https://github.com/keirokeer/OpenUtau-DiffSinger-Lunai)* > - *Singer model from [yousa-ling-official-production/yousa-ling-diffsinger-v1](https://github.com/yousa-ling-official-production/yousa-ling-diffsinger-v1)* diff --git a/README.md b/README.md index 5b98713..b5a8b5e 100644 --- a/README.md +++ b/README.md @@ -17,9 +17,7 @@ * `Pitch Deviation (curve)` * `Tension (curve)` -

- -

+https://github.com/user-attachments/assets/4b5b7c15-947a-4f54-b80e-a14a9eefc86b > - *OpenUtau 版本来自 [keirokeer/OpenUtau-DiffSinger-Lunai](https://github.com/keirokeer/OpenUtau-DiffSinger-Lunai)* > - *歌手模型来自 [yousa-ling-official-production/yousa-ling-diffsinger-v1](https://github.com/yousa-ling-official-production/yousa-ling-diffsinger-v1)* diff --git a/build/auto-py-to-exe.json b/build/auto-py-to-exe.json index 02328a5..c500bad 100644 --- a/build/auto-py-to-exe.json +++ b/build/auto-py-to-exe.json @@ -89,6 +89,10 @@ "optionDest": "datas", "value": "locales;locales/" }, + { + "optionDest": "datas", + "value": "static;static/" + }, { "optionDest": "datas", "value": "README.md;./" diff --git a/build/hatch_build.py b/build/hatch_build.py index ccca363..8b1feae 100644 --- a/build/hatch_build.py +++ b/build/hatch_build.py @@ -1,9 +1,10 @@ -"""Hatch build hook: compile gettext .po -> .mo before wheel packaging.""" +"""Hatch build hook: compile gettext .po -> .mo and download vendored static deps.""" from __future__ import annotations import glob import os +import urllib.request from babel.messages.mofile import write_mo from babel.messages.pofile import read_po @@ -11,9 +12,11 @@ class CustomBuildHook(BuildHookInterface): - PLUGIN_NAME = "custom" - def initialize(self, version: str, build_data: dict) -> None: + self._compile_locales(build_data) + self._vendor_static(build_data) + + def _compile_locales(self, build_data: dict) -> None: locales_dir = os.path.join(self.root, "locales") for po_file in glob.glob( os.path.join(locales_dir, "**", "*.po"), recursive=True @@ -23,7 +26,20 @@ def initialize(self, version: str, build_data: dict) -> None: catalog = read_po(f) with open(mo_file, "wb") as f: write_mo(f, catalog) - # artifacts bypasses .gitignore so the compiled .mo is included in the wheel build_data["artifacts"].append( os.path.relpath(mo_file, self.root) ) + + def _vendor_static(self, build_data: dict) -> None: + vendor_dir = os.path.join(self.root, "static", "vendor") + os.makedirs(vendor_dir, exist_ok=True) + + for name, url in self.config.get("vendor-static-deps", []): + dest = os.path.join(vendor_dir, name) + os.makedirs(os.path.dirname(dest), exist_ok=True) + if not os.path.exists(dest): + print(f"Downloading {name} from {url}") + urllib.request.urlretrieve(url, dest) + build_data["artifacts"].append( + os.path.relpath(dest, self.root) + ) diff --git a/expressions/base.py b/expressions/base.py index cf824d0..7817fb0 100644 --- a/expressions/base.py +++ b/expressions/base.py @@ -7,6 +7,7 @@ import numpy as np from utils.i18n import _, _l +from utils.wavtool import ClampedWav, sec2timestamp from utils.ustx import load_ustx, save_ustx, edit_ustx_expression_curve @@ -25,17 +26,23 @@ class ExpressionLoader(): expression_info: str = "" ustx_lock = threading.Lock() args = SimpleNamespace( - ref_path = Args(name="ref_path" , type=str, default="", help=_l("Path to the **reference** audio file")), # noqa: E501 - utau_path = Args(name="utau_path" , type=str, default="", help=_l("Path to the **UTAU** audio file")), # noqa: E501 - ustx_path = Args(name="ustx_path" , type=str, default="", help=_l("Path to the `.ustx` project file to be processed")), # noqa: E501 - track_number = Args(name="track_number", type=int, default=1 , help=_l("**Track number** to apply expressions to (1-based index)")), # noqa: E501 + ref_path = Args(name="ref_path" , type=str, default="" , help=_l("Path to the **reference** audio file")), # noqa: E501 + utau_path = Args(name="utau_path" , type=str, default="" , help=_l("Path to the **UTAU** audio file")), # noqa: E501 + ustx_path = Args(name="ustx_path" , type=str, default="" , help=_l("Path to the `.ustx` project file to be processed")), # noqa: E501 + track_number = Args(name="track_number", type=int, default=1 , help=_l("**Track number** to apply expressions to (1-based index)")), # noqa: E501 + ref_start = Args(name="ref_start" , type=str, default=None, help=_l("**Start time** of the **reference** audio (format `M:S`, e.g. `0:10.01`). Omit to specify the beginning")), # noqa: E501 + ref_end = Args(name="ref_end" , type=str, default=None, help=_l("**End time** of the **reference** audio (format `M:S`, e.g. `0:10.01`). Omit to specify the ending")), # noqa: E501 + utau_start = Args(name="utau_start" , type=str, default=None, help=_l("**Start time** of the **UTAU** audio (format `M:S`, e.g. `0:10.01`). Omit to specify the beginning")), # noqa: E501 + utau_end = Args(name="utau_end" , type=str, default=None, help=_l("**End time** of the **UTAU** audio (format `M:S`, e.g. `0:10.01`). Omit to specify the ending")), # noqa: E501 ) @classmethod def get_args_dict(cls) -> dict[str, Args]: return cls.args.__dict__ - def __init__(self, ref_path: str, utau_path: str, ustx_path: str): + def __init__(self, ref_path: str, utau_path: str, ustx_path: str, + ref_start: str | None = None, ref_end: str | None = None, + utau_start: str | None = None, utau_end: str | None = None): ExpressionLoader._id_counter += 1 self.id = ExpressionLoader._id_counter self.logger = logging.getLogger(f"{ExpressionLoader.__name__}.{self.expression_name}.{self.id}") @@ -44,8 +51,23 @@ def __init__(self, ref_path: str, utau_path: str, ustx_path: str): self.expression_tick: list | np.ndarray = [] self.expression_val: list | np.ndarray = [] - self.ref_path = ref_path - self.utau_path = utau_path + + self._clamped_ref = ClampedWav(ref_path, ref_start, ref_end, logger=self.logger) + self.ref_path, self.ref_offset, self.ref_duration = ( + self._clamped_ref.path, self._clamped_ref.offset_sec, self._clamped_ref.duration_sec) + self.logger.info(_("ref [{} → {}] {:.3f}s").format( + sec2timestamp(self.ref_offset), + sec2timestamp(self.ref_offset + self.ref_duration), + self.ref_duration)) + + self._clamped_utau = ClampedWav(utau_path, utau_start, utau_end, logger=self.logger) + self.utau_path, self.utau_offset, self.utau_duration = ( + self._clamped_utau.path, self._clamped_utau.offset_sec, self._clamped_utau.duration_sec) + self.logger.info(_("utau [{} → {}] {:.3f}s").format( + sec2timestamp(self.utau_offset), + sec2timestamp(self.utau_offset + self.utau_duration), + self.utau_duration)) + self.ustx_path = ustx_path self.tempo = load_ustx(self.ustx_path)["tempos"][0]["bpm"] self.logger.info(_("Initialization complete.")) diff --git a/expressions/dyn.py b/expressions/dyn.py index 2a7a3e9..4ced865 100644 --- a/expressions/dyn.py +++ b/expressions/dyn.py @@ -1,16 +1,18 @@ from types import SimpleNamespace -import librosa +import numpy as np from scipy.stats import zscore from .base import Args, ExpressionLoader, register_expression from utils.i18n import _, _l from utils.seqtool import ( + time_to_ticks, unify_sequence_time, align_sequence_tick, gaussian_filter1d_with_nan, seq_dynamics_trends, ) +from utils.wavtool import extract_wav_rms @register_expression @@ -18,13 +20,15 @@ class DynLoader(ExpressionLoader): expression_name = "dyn" expression_info = _l("Dynamics (curve)") args = SimpleNamespace( - align_radius = Args(name="align_radius", type=int , default=1 , help=_l("**Radius** for the FastDTW alignment algorithm; larger values allow more flexible alignment but increase computation time")), # noqa: E501 - smoothness = Args(name="smoothness" , type=int , default=2 , help=_l("Controls the **smoothness** of the expression curve using Gaussian filtering. Higher values produce smoother curves but may lose fine detail")), # noqa: E501 - scaler = Args(name="scaler" , type=float, default=1.5, help=_l("**Scaling factor** applied to the expression curve. Values >1 amplify the expression, =1 keeps original intensity, <1 reduces it")), # noqa: E501 + trim_silence = Args(name="trim_silence", type=bool , default=True, help=_l("**Trim silence** from the leading and trailing edges of the audio before extracting expression")), # noqa: E501 + align_radius = Args(name="align_radius", type=int , default=1 , help=_l("**Radius** for the FastDTW alignment algorithm; larger values allow more flexible alignment but increase computation time")), # noqa: E501 + smoothness = Args(name="smoothness" , type=int , default=2 , help=_l("Controls the **smoothness** of the expression curve using Gaussian filtering. Higher values produce smoother curves but may lose fine detail")), # noqa: E501 + scaler = Args(name="scaler" , type=float, default=1.5 , help=_l("**Scaling factor** applied to the expression curve. Values >1 amplify the expression, =1 keeps original intensity, <1 reduces it")), # noqa: E501 ) def get_expression( self, + trim_silence = args.trim_silence.default, align_radius = args.align_radius.default, smoothness = args.smoothness .default, scaler = args.scaler .default, @@ -33,15 +37,15 @@ def get_expression( # Extract rms features from WAV files utau_time, utau_rms, utau_features = get_wav_features( - wav_path=self.utau_path, + wav_path=self.utau_path, mask_silence=trim_silence ) ref_time, ref_rms, ref_features = get_wav_features( - wav_path=self.ref_path, + wav_path=self.ref_path, mask_silence=trim_silence ) # Align all sequences to a common MIDI tick time base # NOTICE: features from UTAU WAV are the reference, and those from Ref. WAV are the query - dyn_tick, (time_aligned_ref_rms, *_unused), *_unused = align_sequence_tick( + dyn_tick, (time_aligned_ref_rms, *_unused), (time_unified_utau_rms, *_unused) = align_sequence_tick( query_time=ref_time, queries=(ref_rms, *ref_features), reference_time=utau_time, @@ -50,27 +54,26 @@ def get_expression( align_radius=align_radius, ) + # Mask positions where utau is silent (NaN) + time_aligned_ref_rms[np.isnan(time_unified_utau_rms)] = np.nan + dyn_val = get_experssion_dynamics(time_aligned_ref_rms, smoothness, scaler) - self.expression_tick, self.expression_val = dyn_tick, dyn_val + # Shift ticks to absolute MIDI position using the UTAU trim offset + utau_offset_ticks = time_to_ticks(self.utau_offset, self.tempo) + self.expression_tick = dyn_tick + utau_offset_ticks + self.expression_val = dyn_val + self.logger.info(_("Expression extraction complete.")) return self.expression_tick, self.expression_val -def extract_wav_rms(wav_path): - sr = librosa.get_samplerate(wav_path) - y, _ = librosa.load(wav_path, sr=sr) - rms = librosa.feature.rms(y=y)[0] - rms_time = librosa.times_like(rms, sr=sr) - return rms_time, rms - - -def get_wav_features(wav_path): +def get_wav_features(wav_path, mask_silence=True): feature_times = [] # List of time sequences(list of lists) feature_vals = [] # List of feature sequences(list of lists) # Extract RMS feature - rms_time, rms = extract_wav_rms(wav_path) + rms_time, rms = extract_wav_rms(wav_path, mask_silence=mask_silence) feature_times += [rms_time] feature_vals += [rms] @@ -89,7 +92,7 @@ def get_wav_features(wav_path): def get_experssion_dynamics(time_aligned_rms, smoothness=2, scaler=1.0): base_scaler = 10.0 smoothed_dyn = gaussian_filter1d_with_nan( - base_scaler * zscore(time_aligned_rms), + base_scaler * zscore(time_aligned_rms, nan_policy='omit'), sigma=smoothness, ) return scaler * smoothed_dyn diff --git a/expressions/pitd.py b/expressions/pitd.py index 596c1a9..5be6614 100644 --- a/expressions/pitd.py +++ b/expressions/pitd.py @@ -1,25 +1,20 @@ -import os -import csv -from pathlib import Path from types import SimpleNamespace -import librosa import numpy as np -from scipy.io import wavfile from scipy.signal import medfilt -from sklearn.decomposition import PCA from librosa import hz_to_midi from .base import Args, ExpressionLoader, register_expression from utils.i18n import _, _l, _lf from utils.seqtool import ( + time_to_ticks, unify_sequence_time, align_sequence_tick, gaussian_filter1d_with_nan, seq_dynamics_trends, ) from utils.log import StreamToLogger -from utils.cache import CACHE_DIR, calculate_file_hash +from utils.wavtool import extract_wav_mfcc, extract_wav_frequency @register_expression @@ -101,44 +96,15 @@ def get_expression( scaler=scaler, ) - self.expression_tick, self.expression_val = pitd_tick, pitd_val + # Shift ticks to absolute MIDI position using the UTAU trim offset + utau_offset_ticks = time_to_ticks(self.utau_offset, self.tempo) + self.expression_tick = pitd_tick + utau_offset_ticks + self.expression_val = pitd_val + self.logger.info(_("Expression extraction complete.")) return self.expression_tick, self.expression_val -def extract_wav_mfcc(wav_path, n_feat=6, n_mfcc=13): - """Extract MFCC features from a WAV file. - - This function extracts Mel-frequency cepstral coefficients (MFCC) from a WAV file. - - Args: - wav_path (str): Path to the WAV file. - n_feat (int, optional): Number of features to extract. Defaults to 6. - n_mfcc (int, optional): Number of MFCC coefficients to extract. Defaults to 13. - - Returns: - tuple: (mfcc_time, mfcc), where: - - mfcc_time (numpy.ndarray): Time points for the MFCC features. Shape: (n_time_points). - - mfcc (numpy.ndarray): Extracted MFCC features. Shape: (n_features, n_time_points). - """ - sr = librosa.get_samplerate(wav_path) - y, _ = librosa.load(wav_path, sr=sr) - - # Extract MFCC features - _mfcc = librosa.feature.mfcc(y=y, sr=sr, n_mfcc=n_mfcc) - mfcc_time = librosa.times_like(_mfcc, sr=sr) - - # Add dynamic features into the MFCC - delta_mfcc = librosa.feature.delta(_mfcc, order=1) - delta2_mfcc = librosa.feature.delta(_mfcc, order=2) - mfcc = np.vstack([_mfcc, delta_mfcc, delta2_mfcc]) - - # PCA to reduce dimensionality - pca = PCA(n_components=n_feat) - mfcc = pca.fit_transform(mfcc.T).T - return mfcc_time, mfcc - - # TODO: Deal with different tempo or ppqn within the same USTX file def get_wav_features(wav_path, backend="swift-f0", confidence_threshold=0.8, confidence_filter_size=9): """Extract features from a WAV file. @@ -236,77 +202,3 @@ def get_pitch_delta(query, reference, scaler=2.5): numpy.ndarray: Scaled pitch difference values. """ return scaler * (query - reference) - - -def extract_wav_frequency(file_path, backend="swift-f0", use_cache=True): - """Extract pitch frequency from a WAV file. - - This function processes an audio file to extract pitch information. - It supports caching to improve performance when processing - the same file multiple times. - - Args: - file_path (str): Path to the WAV file. - backend (str, optional): Pitch detection backend. One of "crepe" or "swift-f0". - "crepe" uses the CREPE model (requires TensorFlow, GPU-accelerated). - "swift-f0" uses SwiftF0 (faster CPU inference, requires swift-f0 package). - Defaults to "swift-f0". - use_cache (bool, optional): Whether to use cached data if available. Defaults to True. - - Returns: - tuple: (time, frequency, confidence), where: - - time (list of float): Time points in seconds. Shape: (n_time_points). - - frequency (list of float): Detected pitch frequencies in Hz. Shape: (n_time_points). - - confidence (list of float): Confidence values for the detected pitches. Shape: (n_time_points). - """ - _SUPPORTED_BACKENDS = ("crepe", "swift-f0") - if backend not in _SUPPORTED_BACKENDS: - raise ValueError(f"Unknown backend '{backend}'. Choose from: {_SUPPORTED_BACKENDS}") - - time = [] - frequency = [] - confidence = [] - cache_dir = Path(CACHE_DIR) / "pitd" - # Try reading data from cache - if use_cache: - os.makedirs(cache_dir, exist_ok=True) - wav_hash = calculate_file_hash(file_path) - - cache_path = cache_dir / f"{wav_hash}.{backend}.csv" - if cache_path.is_file(): - print(_("Loading F0 data from cache file: '{}'").format(cache_path)) - with open(cache_path, "r", newline="") as file: - reader = csv.reader(file) - next(reader) # Skip header - for row in reader: - time.append(float(row[0])) - frequency.append(float(row[1])) - confidence.append(float(row[2])) - - # If cache is unavailable - if not all([time, frequency, confidence]): - # Extract pitch using the specified backend - if backend == "crepe": - import crepe - from utils.gpu import add_cuda_to_path - add_cuda_to_path(skip_missing=True) - sr, audio = wavfile.read(file_path) - time, frequency, confidence, _unused = crepe.predict(audio, sr, viterbi=True) - elif backend == "swift-f0": - from swift_f0 import SwiftF0 - detector = SwiftF0(confidence_threshold=0.0) - result = detector.detect_from_file(file_path) - time = result.timestamps.tolist() - frequency = result.pitch_hz.tolist() - confidence = result.confidence.tolist() - - # Save data to cache - if use_cache: - with open(cache_path, mode="w+", newline="") as file: - writer = csv.writer(file) - writer.writerow(["Time (s)", "Frequency (Hz)", "Confidence"]) - for t, f, c in zip(time, frequency, confidence, strict=False): - writer.writerow([t, f, c]) - print(_("F0 data saved to cache file: '{}'").format(cache_path)) - - return time, frequency, confidence diff --git a/expressions/tenc.py b/expressions/tenc.py index 5078223..e0fea4c 100644 --- a/expressions/tenc.py +++ b/expressions/tenc.py @@ -1,16 +1,18 @@ from types import SimpleNamespace -import librosa +import numpy as np from scipy.stats import zscore from .base import Args, ExpressionLoader, register_expression from utils.i18n import _, _l from utils.seqtool import ( + time_to_ticks, unify_sequence_time, align_sequence_tick, gaussian_filter1d_with_nan, seq_dynamics_trends, ) +from utils.wavtool import extract_wav_rms @register_expression @@ -18,14 +20,16 @@ class TencLoader(ExpressionLoader): expression_name = "tenc" expression_info = _l("Tension (curve)") args = SimpleNamespace( - align_radius = Args(name="align_radius", type=int , default=1 , help=_l("**Radius** for the FastDTW alignment algorithm; larger values allow more flexible alignment but increase computation time")), # noqa: E501 - smoothness = Args(name="smoothness" , type=int , default=6 , help=_l("Controls the **smoothness** of the expression curve using Gaussian filtering. Higher values produce smoother curves but may lose fine detail")), # noqa: E501 - scaler = Args(name="scaler" , type=float, default=1.0, help=_l("**Scaling factor** applied to the expression curve. Values >1 amplify the expression, =1 keeps original intensity, <1 reduces it")), # noqa: E501 - bias = Args(name="bias" , type=int , default=10 , help=_l("**Bias** offset added to the expression curve. Positive values shift the curve upward; negative values shift it downward")), # noqa: E501 + trim_silence = Args(name="trim_silence", type=bool , default=True, help=_l("**Trim silence** from the leading and trailing edges of the audio before extracting expression")), # noqa: E501 + align_radius = Args(name="align_radius", type=int , default=1 , help=_l("**Radius** for the FastDTW alignment algorithm; larger values allow more flexible alignment but increase computation time")), # noqa: E501 + smoothness = Args(name="smoothness" , type=int , default=6 , help=_l("Controls the **smoothness** of the expression curve using Gaussian filtering. Higher values produce smoother curves but may lose fine detail")), # noqa: E501 + scaler = Args(name="scaler" , type=float, default=1.0 , help=_l("**Scaling factor** applied to the expression curve. Values >1 amplify the expression, =1 keeps original intensity, <1 reduces it")), # noqa: E501 + bias = Args(name="bias" , type=int , default=10 , help=_l("**Bias** offset added to the expression curve. Positive values shift the curve upward; negative values shift it downward")), # noqa: E501 ) def get_expression( self, + trim_silence = args.trim_silence.default, align_radius = args.align_radius.default, smoothness = args.smoothness .default, scaler = args.scaler .default, @@ -35,15 +39,15 @@ def get_expression( # Extract features from WAV files utau_time, utau_rms, utau_features = get_wav_features( - wav_path=self.utau_path, + wav_path=self.utau_path, mask_silence=trim_silence ) ref_time, ref_rms, ref_features = get_wav_features( - wav_path=self.ref_path, + wav_path=self.ref_path, mask_silence=trim_silence ) # Align all sequences to a common MIDI tick time base # NOTICE: features from UTAU WAV are the reference, and those from Ref. WAV are the query - tenc_tick, (time_aligned_ref_rms, *_unused), *_unused = align_sequence_tick( + tenc_tick, (time_aligned_ref_rms, *_unused), (time_unified_utau_rms, *_unused) = align_sequence_tick( query_time=ref_time, queries=(ref_rms, *ref_features), reference_time=utau_time, @@ -52,27 +56,26 @@ def get_expression( align_radius=align_radius, ) + # Mask positions where utau is silent (NaN) + time_aligned_ref_rms[np.isnan(time_unified_utau_rms)] = np.nan + tenc_val = get_experssion_tension(time_aligned_ref_rms, smoothness, scaler, bias) - self.expression_tick, self.expression_val = tenc_tick, tenc_val + # Shift ticks to absolute MIDI position using the UTAU trim offset + utau_offset_ticks = time_to_ticks(self.utau_offset, self.tempo) + self.expression_tick = tenc_tick + utau_offset_ticks + self.expression_val = tenc_val + self.logger.info(_("Expression extraction complete.")) return self.expression_tick, self.expression_val -def extract_wav_rms(wav_path): - sr = librosa.get_samplerate(wav_path) - y, _ = librosa.load(wav_path, sr=sr) - rms = librosa.feature.rms(y=y)[0] - rms_time = librosa.times_like(rms, sr=sr) - return rms_time, rms - - -def get_wav_features(wav_path): +def get_wav_features(wav_path, mask_silence=True): feature_times = [] # List of time sequences(list of lists) feature_vals = [] # List of feature sequences(list of lists) # Extract RMS - rms_time, rms = extract_wav_rms(wav_path) + rms_time, rms = extract_wav_rms(wav_path, mask_silence=mask_silence) feature_times += [rms_time] feature_vals += [rms] @@ -91,7 +94,7 @@ def get_wav_features(wav_path): def get_experssion_tension(time_aligned_rms, smoothness=2, scaler=1.0, bias=0): base_scaler = 10.0 smoothed_tenc = gaussian_filter1d_with_nan( - base_scaler * zscore(time_aligned_rms), + base_scaler * zscore(time_aligned_rms, nan_policy='omit'), sigma=smoothness, ) return scaler * smoothed_tenc + bias diff --git a/expressive.py b/expressive.py index 7ee1b3e..119df5f 100644 --- a/expressive.py +++ b/expressive.py @@ -7,7 +7,10 @@ from os.path import splitext, basename from __version__ import VERSION -from utils.cli import ArgumentDefaultsWrappedTextRichHelpFormatter +from utils.cli import ( + add_expression_args_group, + ArgumentDefaultsWrappedTextRichHelpFormatter, +) from expressions.base import getExpressionLoader, get_registered_expressions @@ -17,6 +20,10 @@ def process_expressions( ustx_input: str, ustx_output: str, track_number: int, + ref_start: str | None, + ref_end: str | None, + utau_start: str | None, + utau_end: str | None, expressions: list[dict], ): """ @@ -28,6 +35,10 @@ def process_expressions( ustx_input (str): Path to the input USTX project file. ustx_output (str): Path to save the processed USTX project file. track_number (int): Track number to apply expressions. + ref_start (str | None): Start time of the reference audio in M:S format. Defaults to None. + ref_end (str | None): End time of the reference audio in M:S format. Defaults to None. + utau_start (str | None): Start time of the UTAU audio in M:S format. Defaults to None. + utau_end (str | None): End time of the UTAU audio in M:S format. Defaults to None. expressions (list[dict]): List of expressions to process, each containing: - "expression": Expression type (e.g., "dyn", "pitd", "tenc"). - Additional parameters specific to the expression type. @@ -67,7 +78,11 @@ def process_expressions( if exp_type not in get_registered_expressions(): raise ValueError(f"Expression '{exp_type}' is not supported.") - loader = getExpressionLoader(exp_type)(ref_wav, utau_wav, ustx_output) + loader = getExpressionLoader(exp_type)( + ref_wav, utau_wav, ustx_output, + ref_start=ref_start, ref_end=ref_end, + utau_start=utau_start, utau_end=utau_end, + ) loader_args = { arg_name: exp.get(arg_name, arg.default) for arg_name, arg in loader.get_args_dict().items() @@ -131,8 +146,12 @@ def main(): parser.add_argument("-i", "--ustx_input", type=general_args.ustx_path.type, required=True, help=general_args.ustx_path.help) # noqa: E501 parser.add_argument("-o", "--ustx_output", type=str, required=True, help="Path to save the processed `.ustx` file") # noqa: E501 parser.add_argument("-t", "--track_number", type=general_args.track_number.type, required=True, help=general_args.track_number.help) # noqa: E501 - - parser.add_argument("-e", "--expression", type=str, action="append", required=True, choices=get_registered_expressions(), + parser.add_argument("--utau_start", type=general_args.utau_start.type, default=general_args.utau_start.default, help=general_args.utau_start.help) # noqa: E501 + parser.add_argument("--utau_end", type=general_args.utau_end.type, default=general_args.utau_end.default, help=general_args.utau_end.help) # noqa: E501 + parser.add_argument("--ref_start", type=general_args.ref_start.type, default=general_args.ref_start.default, help=general_args.ref_start.help) # noqa: E501 + parser.add_argument("--ref_end", type=general_args.ref_end.type, default=general_args.ref_end.default, help=general_args.ref_end.help) # noqa: E501 + + parser.add_argument("-e", "--expression", type=str, action="append", required=True, choices=get_registered_expressions(), help="**Expression(s)** to apply. Repeat the flag for multiple expressions (e.g., `-e dyn -e pitd`)") parser.add_argument("--version", action="version", version=f"%(prog)s v{VERSION}") @@ -141,12 +160,7 @@ def main(): get_expression_args = lambda exp_name: getExpressionLoader(exp_name).get_args_dict() for exp_name in expression_names: - exp_info = getExpressionLoader(exp_name).expression_info - group = parser.add_argument_group(f"[{exp_name.upper()}] {exp_info} Expression") - for arg_name, arg in get_expression_args(exp_name).items(): - group.add_argument(f"--{exp_name}.{arg_name}", - type=arg.type, default=arg.default, help=arg.help, - choices=arg.choices) + add_expression_args_group(parser, exp_name) # Parse arguments args = parser.parse_args() @@ -166,7 +180,10 @@ def main(): try: process_expressions( args.utau_wav, args.ref_wav, args.ustx_input, - args.ustx_output, args.track_number, expressions + args.ustx_output, args.track_number, + args.ref_start, args.ref_end, + args.utau_start, args.utau_end, + expressions, ) except Exception as e: logger_app.error(f"Error occurred during processing: {e}") diff --git a/expressive_gui.py b/expressive_gui.py index 18e9978..ea427ef 100644 --- a/expressive_gui.py +++ b/expressive_gui.py @@ -1,6 +1,7 @@ import os import sys import json +import time import logging import asyncio import argparse @@ -11,12 +12,13 @@ from concurrent.futures import ProcessPoolExecutor import webview -from nicegui import ui, app +from nicegui import ui, app, background_tasks from utils.ui import ( blink_taskbar_window, change_window_style, NiceguiNativeDropArea, + WaveSurferRangeSelector, webview_active_window, ) from utils.monkeypatch import ( @@ -26,8 +28,9 @@ ) from __version__ import VERSION from utils.i18n import _, init_gettext -from utils.worker import WorkerContext, setup_worker_context from expressive import process_expressions +from utils.wavtool import get_wav_end_ts, validate_timestamp +from utils.worker import WorkerContext, setup_worker_context from expressions.base import getExpressionLoader, get_registered_expressions @@ -51,6 +54,8 @@ domain=LOCALE_DOMAIN, ) +general_args = getExpressionLoader(None).args + class LogElementHandler(logging.Handler): """A logging handler that emits messages to a log element.""" @@ -63,6 +68,7 @@ def emit(self, record: logging.LogRecord) -> None: try: msg = self.format(record) self.element.push(msg) + time.sleep(0.1) # Avoid flooding except Exception: self.handleError(record) @@ -104,7 +110,7 @@ def dict_update(d: dict, u: Mapping): def close_splash(): """Close the splash screen when the app is connected if this script is frozen. This is a workaround for PyInstaller, which doesn't support splash screen in the main thread - + See: https://github.com/zauberzeug/nicegui/discussions/3536 https://stackoverflow.com/questions/71057636/how-can-i-solve-no-module-named-pyi-splash-after-using-pyinstaller """ @@ -113,14 +119,18 @@ def close_splash(): pyi_splash.close() -def create_gui(): - # Initialize state - state = { - "utau_wav" : "", - "ref_wav" : "", - "ustx_input" : "", +def build_default_state() -> dict: + """Build the default application state from registered expression loaders.""" + return { + "utau_wav" : general_args.utau_path.default, + "ref_wav" : general_args.ref_path.default, + "ustx_input" : general_args.ustx_path.default, "ustx_output" : "", - "track_number": 1, + "track_number": general_args.track_number.default, + "ref_start" : general_args.ref_start.default, + "ref_end" : general_args.ref_end.default, + "utau_start" : general_args.utau_start.default, + "utau_end" : general_args.utau_end.default, "expressions" : { exp_name: { "selected": False, @@ -132,30 +142,9 @@ def create_gui(): }, } - # state = { - # "utau_wav" : "", - # "ref_wav" : "", - # "ustx_input" : "", - # "ustx_output" : "", - # "track_number": 1, - # "expressions" : { - # "dyn": { - # "selected" : False, - # "align_radius": 1, - # "smoothness" : 2, - # "scaler" : 2.0, - # }, - # "pitd": { - # "selected" : False, - # "confidence_utau": 0.8, - # "confidence_ref" : 0.6, - # "align_radius" : 1, - # "semitone_shift" : None, - # "smoothness" : 2, - # "scaler" : 2.0, - # }, - # }, - # } + +def create_gui(): # noqa: C901 + state = build_default_state() def on_color_scheme_changed(event): """Change the window style based on the color scheme.""" @@ -188,7 +177,10 @@ async def import_config(state=state): try: with open(file[0], "r", encoding="utf-8-sig") as f: cfg = json.load(f) - dict_update(state, cfg) + # Start from defaults, then overlay imported values — missing keys stay default + default_state = build_default_state() + dict_update(default_state, cfg) + dict_update(state, default_state) ui.notify(_("Config imported successfully!"), type="positive") ui.update() @@ -235,6 +227,10 @@ async def run_processing(): state["ustx_input"], state["ustx_output"], state["track_number"], + state["ref_start"], + state["ref_end"], + state["utau_start"], + state["utau_end"], expressions, ) blink_taskbar_window(app.config.title) @@ -343,7 +339,7 @@ def configure_logger(name: str, formatter: logging.Formatter): +""") # noqa: E501 + + +# --------------------------------------------------------------------------- +# WaveSurferRangeSelector +# --------------------------------------------------------------------------- + +def seconds_to_timestamp(seconds: float) -> str: + """Convert seconds to 'm:ss.ss' format matching ExpressionLoader timestamp style.""" + m = int(seconds) // 60 + s = seconds - m * 60 + return f"{m}:{s:05.2f}" + + +def serve_wav(wav_path: str) -> str: + """Serve a local WAV file via NiceGUI's static file server and return its URL. + + Each unique directory is registered once under /wav/. + WaveSurfer then streams the file normally — no base64 overhead. + """ + import hashlib + directory = os.path.dirname(os.path.abspath(wav_path)) + dir_hash = hashlib.md5(directory.encode()).hexdigest()[:8] + mount = f"/wav/{dir_hash}" + # Register the directory only once + if not any(r.path == mount for r in app.routes): + app.add_static_files(mount, directory) + filename = os.path.basename(wav_path) + return f"{mount}/{filename}" + + +class WaveSurferRangeSelector(ui.element): + """WaveSurferElement with NiceGUI-style BindableProperty bindings. + + - ``wav_path``: one-way in via ``bind_wav_path_from(obj, key)`` — reloads on change. + - ``start`` / ``end``: two-way via ``bind_start`` / ``bind_end`` — region drag + writes back to the bound object in 'm:ss.ss' format. + + Usage:: + + ws = WaveSurferRangeSelector() + ws.bind_wav_path_from(state, 'ref_wav') + ws.bind_start(state, 'ref_start') + ws.bind_end(state, 'ref_end') + """ + + from nicegui.binding import BindableProperty + + wav_path = BindableProperty( + on_change=lambda self, val: self._on_wav_path_change(val)) # type: ignore[misc] + start = BindableProperty( + on_change=lambda self, val: self._on_range_change()) # type: ignore[misc] + end = BindableProperty( + on_change=lambda self, val: self._on_range_change()) # type: ignore[misc] + + def __init__( + self, + *, + wave_color: str = "rgb(200, 0, 200)", + progress_color: str = "rgb(100, 0, 100)", + height: int = 80, + bar_width: int = 2, + bar_gap: int = 1, + bar_radius: int = 2, + ) -> None: + super().__init__(tag="div") + self.wav_path: str = "" + self.start: str = "" + self.end: str = "" + self._from_js: bool = False # guard against JS→Python→JS round-trip + + self.classes("w-full") + + with self: + self._ws = WaveSurferElement( + wave_color=wave_color, + progress_color=progress_color, + height=height, + bar_width=bar_width, + bar_gap=bar_gap, + bar_radius=bar_radius, + loop_regions=False, + enable_drag_selection=True, + drag_selection_color="rgba(100,180,255,0.25)", + show_controls=True, + ) + + iid = self._ws._iid + ui.add_body_html(f""" + +""") + ui.on(f'wavesurfer-range-{iid}', self._handle_region_updated) + ui.on(f"{iid}-ready", self._on_range_change) + + def _handle_region_updated(self, e) -> None: + """Called by JS on region-updated; writes back through BindableProperty.""" + self._from_js = True + self.start = e.args['start'] or None + self.end = e.args['end'] or None + self._from_js = False + + @staticmethod + def _timestamp_to_seconds(ts: str | None) -> float | None: + """Parse 'm:ss.ss' → seconds, or bare float string → seconds. Returns None if invalid/empty.""" + if not ts: + return None + try: + if ':' in ts: + m, s = ts.split(':', 1) + return int(m) * 60 + float(s) + return float(ts) + except ValueError: + return None + + def _on_range_change(self) -> None: + """Called by BindableProperty when start or end changes from the Python side.""" + if self._from_js: + return # already came from JS, don't echo back + # Guard: binding may fire before NiceGUI's event loop is ready (e.g. at startup) + from nicegui import core + if core.loop is None: + return + + iid = self._ws._iid + start_s = self._timestamp_to_seconds(self.start or '') + end_s = self._timestamp_to_seconds(self.end or '') + + # Invalid parse (non-empty but unparseable) or start >= end → clear + start_invalid = bool((self.start or '').strip()) and start_s is None + end_invalid = bool((self.end or '').strip()) and end_s is None + if start_invalid or end_invalid or ( + start_s is not None and end_s is not None and start_s >= end_s + ): + self._ws._js_client.run_javascript( + f"window['{self._ws._iid}']?.regions.clearRegions()" + ) + return + + # Empty start = beginning (0), empty end = track duration + start_js = start_s if start_s is not None else 0 + end_js = f'{end_s}' if end_s is not None else 'inst.ws.getDuration()' + self._ws._js_client.run_javascript(f""" + (function applyRegion() {{ + const inst = window['{iid}']; + if (!inst) {{ setTimeout(applyRegion, 80); return; }} + const apply = () => {{ + const start = {start_js}; + const end = {end_js}; + if (start >= end) return; + inst._updatingFromPython = true; + const regions = inst.regions.getRegions(); + if (regions.length > 0) {{ + regions[0].setOptions({{ start, end }}); + }} else {{ + inst.regions.addRegion({{ + start, end, + color: 'rgba(100,180,255,0.25)', + drag: true, resize: true, + }}); + }} + inst._updatingFromPython = false; + }}; + if (inst.ws.getDuration() > 0.001) {{ apply(); }} + }})(); + """) + + def _on_wav_path_change(self, path: str) -> None: + """Called automatically by BindableProperty when wav_path changes. + Clears the waveform, then loads the new file. + Re-renders the region once the waveform is ready. + """ + # Guard: binding may fire before NiceGUI's event loop is ready (e.g. at startup) + from nicegui import core + if core.loop is None: + return + if path: + self._ws.empty() + self._ws.clear_regions() + self._ws.zoom(0) + self._ws.load(serve_wav(path)) + + def bind_wav_path_from(self, target_object: Any, target_name: str = 'wav_path') -> "WaveSurferRangeSelector": + """One-way bind: target → wav_path.""" + from nicegui.binding import bind_from + bind_from(self, 'wav_path', target_object, target_name) + return self + + def bind_start(self, target_object: Any, target_name: str = 'start') -> "WaveSurferRangeSelector": + """Two-way bind: start ↔ target.""" + from nicegui.binding import bind + bind(self, 'start', target_object, target_name) + return self + + def bind_end(self, target_object: Any, target_name: str = 'end') -> "WaveSurferRangeSelector": + """Two-way bind: end ↔ target.""" + from nicegui.binding import bind + bind(self, 'end', target_object, target_name) + return self + + +# --------------------------------------------------------------------------- +# Misc helpers +# --------------------------------------------------------------------------- + def tooltip_md(element: ui.element, text: str) -> ui.element: """Add a markdown-rendered tooltip to a NiceGUI element. Chainable like .tooltip().""" with element: diff --git a/utils/wavtool.py b/utils/wavtool.py new file mode 100644 index 0000000..3f193ad --- /dev/null +++ b/utils/wavtool.py @@ -0,0 +1,346 @@ +import os +import csv +import atexit +import logging +import argparse +import tempfile +from pathlib import Path + +import librosa +import numpy as np +import soundfile as sf +from scipy.io import wavfile +from sklearn.decomposition import PCA +from skimage.filters import threshold_otsu + +from utils.i18n import _ +from utils.cache import CACHE_DIR, calculate_file_hash + + +def extract_wav_mfcc(wav_path, n_feat=6, n_mfcc=13): + """Extract MFCC features from a WAV file. + + This function extracts Mel-frequency cepstral coefficients (MFCC) from a WAV file. + + Args: + wav_path (str): Path to the WAV file. + n_feat (int, optional): Number of features to extract. Defaults to 6. + n_mfcc (int, optional): Number of MFCC coefficients to extract. Defaults to 13. + + Returns: + tuple: (mfcc_time, mfcc), where: + - mfcc_time (numpy.ndarray): Time points for the MFCC features. Shape: (n_time_points). + - mfcc (numpy.ndarray): Extracted MFCC features. Shape: (n_features, n_time_points). + """ + sr = librosa.get_samplerate(wav_path) + y, _ = librosa.load(wav_path, sr=sr) + + # Extract MFCC features + _mfcc = librosa.feature.mfcc(y=y, sr=sr, n_mfcc=n_mfcc) + mfcc_time = librosa.times_like(_mfcc, sr=sr) + + # Add dynamic features into the MFCC + delta_mfcc = librosa.feature.delta(_mfcc, order=1) + delta2_mfcc = librosa.feature.delta(_mfcc, order=2) + mfcc = np.vstack([_mfcc, delta_mfcc, delta2_mfcc]) + + # PCA to reduce dimensionality + pca = PCA(n_components=n_feat) + mfcc = pca.fit_transform(mfcc.T).T + return mfcc_time, mfcc + + +def extract_wav_frequency(file_path, backend="swift-f0", use_cache=True): + """Extract pitch frequency from a WAV file. + + This function processes an audio file to extract pitch information. + It supports caching to improve performance when processing + the same file multiple times. + + Args: + file_path (str): Path to the WAV file. + backend (str, optional): Pitch detection backend. One of "crepe" or "swift-f0". + "crepe" uses the CREPE model (requires TensorFlow, GPU-accelerated). + "swift-f0" uses SwiftF0 (faster CPU inference, requires swift-f0 package). + Defaults to "swift-f0". + use_cache (bool, optional): Whether to use cached data if available. Defaults to True. + + Returns: + tuple: (time, frequency, confidence), where: + - time (list of float): Time points in seconds. Shape: (n_time_points). + - frequency (list of float): Detected pitch frequencies in Hz. Shape: (n_time_points). + - confidence (list of float): Confidence values for the detected pitches. Shape: (n_time_points). + """ + _SUPPORTED_BACKENDS = ("crepe", "swift-f0") + if backend not in _SUPPORTED_BACKENDS: + raise ValueError(f"Unknown backend '{backend}'. Choose from: {_SUPPORTED_BACKENDS}") + + time = [] + frequency = [] + confidence = [] + cache_dir = Path(CACHE_DIR) / "pitd" + # Try reading data from cache + if use_cache: + os.makedirs(cache_dir, exist_ok=True) + wav_hash = calculate_file_hash(file_path) + + cache_path = cache_dir / f"{wav_hash}.{backend}.csv" + if cache_path.is_file(): + print(_("Loading F0 data from cache file: '{}'").format(cache_path)) + with open(cache_path, "r", newline="") as file: + reader = csv.reader(file) + next(reader) # Skip header + for row in reader: + time.append(float(row[0])) + frequency.append(float(row[1])) + confidence.append(float(row[2])) + + # If cache is unavailable + if not all([time, frequency, confidence]): + # Extract pitch using the specified backend + if backend == "crepe": + import crepe + from utils.gpu import add_cuda_to_path + add_cuda_to_path(skip_missing=True) + sr, audio = wavfile.read(file_path) + time, frequency, confidence, _unused = crepe.predict(audio, sr, viterbi=True) + elif backend == "swift-f0": + from swift_f0 import SwiftF0 + detector = SwiftF0(confidence_threshold=0.0) + result = detector.detect_from_file(file_path) + time = result.timestamps.tolist() + frequency = result.pitch_hz.tolist() + confidence = result.confidence.tolist() + + # Save data to cache + if use_cache: + with open(cache_path, mode="w+", newline="") as file: + writer = csv.writer(file) + writer.writerow(["Time (s)", "Frequency (Hz)", "Confidence"]) + for t, f, c in zip(time, frequency, confidence, strict=False): + writer.writerow([t, f, c]) + print(_("F0 data saved to cache file: '{}'").format(cache_path)) + + return time, frequency, confidence + + +def extract_wav_rms(wav_path, mask_silence=True): + """Extract RMS energy from a WAV file. + + Args: + wav_path (str): Path to the WAV file. + mask_silence (bool, optional): If True, masks leading and trailing silence with NaN + using Otsu's method to auto-detect the silence threshold. Defaults to True. + + Returns: + tuple: (rms_time, rms), where: + - rms_time (numpy.ndarray): Time values for each RMS frame. Shape: (n_frames,). + - rms (numpy.ndarray): RMS energy values, with NaN at silent edges if mask_silence + is True. Shape: (n_frames,). + """ + sr = librosa.get_samplerate(wav_path) + y, _ = librosa.load(wav_path, sr=sr) + rms = librosa.feature.rms(y=y)[0] + rms_time = librosa.times_like(rms, sr=sr) + if mask_silence: + threshold = threshold_otsu(rms) + is_silent = rms < threshold + start_frame = np.argmax(~is_silent) + end_frame = len(is_silent) - np.argmax(~is_silent[::-1]) + rms[:start_frame] = np.nan + rms[end_frame:] = np.nan + return rms_time, rms + + +def timestamp2sec(value: str) -> float: + """Parse a timestamp string in M:S format (e.g. '0:10.01') into seconds. + + Intended for use as ``type=timestamp2sec`` in + :func:`argparse.ArgumentParser.add_argument`, so argparse stores the + result directly as a ``float`` number of seconds. + + Args: + value (str): The timestamp string to parse. + + Returns: + float: Total time in seconds (e.g. '1:30.5' -> ``90.5``). + + Raises: + argparse.ArgumentTypeError: If the string is not a valid M:S timestamp. + """ + parts = value.split(":") + if len(parts) != 2: + raise argparse.ArgumentTypeError( + f"Invalid timestamp '{value}'. Expected M:S (e.g. '0:10.01')." + ) + minutes_str, seconds_str = parts + try: + minutes = int(minutes_str) + seconds = float(seconds_str) + except ValueError as err: + raise argparse.ArgumentTypeError( + f"Invalid timestamp '{value}'. " + "Minutes must be an integer and seconds must be a number (e.g. '0:10.01')." + ) from err + if minutes < 0: + raise argparse.ArgumentTypeError( + f"Invalid timestamp '{value}': minutes must be non-negative, got {minutes}." + ) + if not (0 <= seconds < 60): + raise argparse.ArgumentTypeError( + f"Invalid timestamp '{value}': seconds must be in [0, 60), got {seconds}." + ) + return minutes * 60.0 + seconds + + +def validate_timestamp(value: str | None, arg_name: str) -> bool: + """Validate a timestamp argument in M:S format (e.g. '0:10.01'). + + Wraps :func:`timestamp2sec` for use outside argparse. + Accepts ``None`` silently (meaning "use default boundary"). + + Args: + value (str | None): The timestamp string to validate, or None to skip. + arg_name (str): The argument name, used in error messages. + + Returns: + bool: ``True`` if *value* is ``None`` or a valid M:S timestamp, + ``False`` otherwise. + """ + if value is None: + return True + try: + timestamp2sec(value) + return True + except argparse.ArgumentTypeError: + return False + + +def sec2timestamp(sec: float) -> str: + """Format seconds as a M:SS.ss timestamp string (e.g. '1:05.30'). + + Args: + sec (float): Time in seconds. + + Returns: + str: Formatted timestamp string. + """ + m = int(sec) // 60 + s = sec - m * 60 + return f"{m}:{s:05.2f}" + + +def get_wav_end_ts(wav_path: str): + return sec2timestamp(librosa.get_duration(path=wav_path)) + + +class ClampedWav: + """Trim a WAV file to [ts_start, ts_end] and manage the resulting temp file. + + The trimmed audio is written to a temporary WAV file on construction. + The temp file is deleted automatically when: + + * the instance is garbage-collected (``__del__``), or + * the Python process exits normally or via an unhandled exception + (``atexit`` handler). + + Use as a plain object **or** as a context manager (``with`` statement) for + deterministic, prompt cleanup: + + .. code-block:: python + + with ClampedWav(wav_path, "0:10", "1:30") as clamped: + process(clamped.path) + # temp file already gone here + + Attributes: + path (str): Path to the temporary trimmed WAV file. + offset_sec (float): Start position inside the original file (seconds). + duration_sec (float): Length of the trimmed segment (seconds). + """ + + def __init__( + self, + wav_path: str, + ts_start: str | None, + ts_end: str | None, + logger: logging.Logger | logging.LoggerAdapter | None = None, + ) -> None: + """Trim *wav_path* to [ts_start, ts_end] and write it to a temp file. + + Both timestamps are clamped to ``[0, duration]`` before trimming. + + Args: + wav_path (str): Path to the source WAV file. + ts_start (str | None): Start timestamp in M:S format, or ``None`` + for the beginning of the file. + ts_end (str | None): End timestamp in M:S format, or ``None`` for + the end of the file. + logger: Optional logger for clamp warnings. + """ + total_duration = librosa.get_duration(path=wav_path) + + start_sec = timestamp2sec(ts_start) if ts_start is not None else 0.0 + end_sec = timestamp2sec(ts_end) if ts_end is not None else total_duration + + # Clamp to valid range + start_clamped = max(0.0, min(start_sec, total_duration)) + end_clamped = max(0.0, min(end_sec, total_duration)) + + if logger is not None: + if start_clamped != start_sec: + logger.warning( + _("start {:.3f}s clamped to {:.3f}s (total duration: {:.3f}s)").format( + start_sec, start_clamped, total_duration + ) + ) + if end_clamped != end_sec: + logger.warning( + _("end {:.3f}s clamped to {:.3f}s (total duration: {:.3f}s)").format( + end_sec, end_clamped, total_duration + ) + ) + + self.offset_sec = start_clamped + self.duration_sec = end_clamped - start_clamped + + # Write trimmed audio to a named temp file + y, sr = librosa.load( + wav_path, sr=None, offset=self.offset_sec, duration=self.duration_sec + ) + tmp = tempfile.NamedTemporaryFile(suffix=".wav", delete=False) + sf.write(tmp.name, y, sr) + tmp.close() + + self.path: str = tmp.name + + # Register atexit so the file is removed even if __del__ is skipped + # (e.g. interpreter shutdown, unhandled exception, or reference cycles). + atexit.register(self._cleanup) + + # ------------------------------------------------------------------ + # Cleanup helpers + # ------------------------------------------------------------------ + + def _cleanup(self) -> None: + """Delete the temp file if it still exists. Safe to call multiple times.""" + path, self.path = getattr(self, "path", None), "" + if path: + try: + os.unlink(path) + except FileNotFoundError: + pass # already gone — that's fine + + def __del__(self) -> None: + self._cleanup() + + # ------------------------------------------------------------------ + # Context-manager support + # ------------------------------------------------------------------ + + def __enter__(self) -> "ClampedWav": + return self + + def __exit__(self, exc_type, exc_val, exc_tb) -> None: + self._cleanup() + return None # do not suppress exceptions