Initial commit

This commit is contained in:
王新升
2026-02-06 20:31:14 +08:00
parent a0b51be095
commit c589bcb837
145 changed files with 28773 additions and 0 deletions
@@ -0,0 +1,129 @@
import numpy as np
import os
import soundfile as sf
import matplotlib.pyplot as plt
from typing import Dict, Tuple, Optional
import torch.distributed as dist
def read_audio_transposed(path: str, instr: Optional[str] = None, skip_err: bool = False) -> Tuple[Optional[np.ndarray], Optional[int]]:
"""
Read an audio file and return transposed waveform data with channels first.
Loads the audio file from `path`, converts mono signals to 2D format, and
transposes the array so that its shape is (channels, length). In case of
errors, either raises an exception or skips gracefully depending on
`skip_err`.
Args:
path (str): Path to the audio file to load.
instr (Optional[str], optional): Instrument name, used for informative
messages when `skip_err` is True. Defaults to None.
skip_err (bool, optional): If True, skip files with read errors and
return `(None, None)` instead of raising. Defaults to False.
Returns:
Tuple[Optional[np.ndarray], Optional[int]]: A tuple containing:
- NumPy array of shape (channels, length), or None if skipped.
- Sampling rate as an integer, or None if skipped.
"""
should_print = not dist.is_initialized() or dist.get_rank() == 0
try:
mix, sr = sf.read(path)
except Exception as e:
if skip_err:
if should_print:
print(f"No stem {instr}: skip!")
return None, None
else:
raise RuntimeError(f"Error reading the file at {path}: {e}")
else:
if len(mix.shape) == 1: # For mono audio
mix = np.expand_dims(mix, axis=-1)
return mix.T, sr
def normalize_audio(audio: np.ndarray) -> Tuple[np.ndarray, Dict[str, float]]:
"""
Normalize an audio signal using mean and standard deviation.
Computes the mean and standard deviation from the mono mix of the input
signal, then applies normalization to each channel.
Args:
audio (np.ndarray): Input audio array of shape (channels, time) or (time,).
Returns:
Tuple[np.ndarray, Dict[str, float]]: A tuple containing:
- Normalized audio with the same shape as the input.
- A dictionary with keys "mean" and "std" from the original audio.
"""
mono = audio.mean(0)
mean, std = mono.mean(), mono.std()
return (audio - mean) / std, {"mean": mean, "std": std}
def denormalize_audio(audio: np.ndarray, norm_params: Dict[str, float]) -> np.ndarray:
"""
Reverse normalization on an audio signal.
Applies the stored mean and standard deviation to restore the original
scale of a previously normalized signal.
Args:
audio (np.ndarray): Normalized audio array to be denormalized.
norm_params (Dict[str, float]): Dictionary containing the keys
"mean" and "std" used during normalization.
Returns:
np.ndarray: Denormalized audio with the same shape as the input.
"""
return audio * norm_params["std"] + norm_params["mean"]
def draw_spectrogram(waveform: np.ndarray, sample_rate: int, length: float, output_file: str) -> None:
"""
Generate and save a spectrogram image from an audio waveform.
Converts the provided waveform into a mono signal, computes its Short-Time
Fourier Transform (STFT), converts the amplitude spectrogram to dB scale,
and plots it using a plasma colormap.
Args:
waveform (np.ndarray): Input audio waveform array of shape (time, channels)
or (time,).
sample_rate (int): Sampling rate of the waveform in Hz.
length (float): Duration (in seconds) of the waveform to include in the
spectrogram.
output_file (str): Path to save the resulting spectrogram image.
Returns:
None
"""
import librosa.display
# Cut only required part of spectorgram
x = waveform[:int(length * sample_rate), :]
X = librosa.stft(x.mean(axis=-1)) # perform short-term fourier transform on mono signal
Xdb = librosa.amplitude_to_db(np.abs(X), ref=np.max) # convert an amplitude spectrogram to dB-scaled spectrogram.
fig, ax = plt.subplots()
# plt.figure(figsize=(30, 10)) # initialize the fig size
img = librosa.display.specshow(
Xdb,
cmap='plasma',
sr=sample_rate,
x_axis='time',
y_axis='linear',
ax=ax
)
ax.set(title='File: ' + os.path.basename(output_file))
fig.colorbar(img, ax=ax, format="%+2.f dB")
if output_file is not None:
plt.savefig(output_file)
@@ -0,0 +1,421 @@
import numpy as np
import torch
import librosa
import torch.nn.functional as F
from typing import Dict, List, Tuple
def sdr(references: np.ndarray, estimates: np.ndarray) -> float:
"""
Compute Signal-to-Distortion Ratio (SDR) for one or more audio tracks.
SDR is a measure of how well the predicted source (estimate) matches the reference source.
It is calculated as the ratio of the energy of the reference signal to the energy of the error (difference between reference and estimate).
Return SDR in decibels (dB)
Parameters:
----------
references : np.ndarray
A 3D numpy array of shape (num_sources, num_channels, num_samples), where num_sources is the number of sources,
num_channels is the number of channels (e.g., 1 for mono, 2 for stereo), and num_samples is the length of the audio signal.
estimates : np.ndarray
A 3D numpy array of shape (num_sources, num_channels, num_samples) representing the estimated sources.
Returns:
-------
np.ndarray
A 1D numpy array containing the SDR values for each source.
"""
eps = 1e-8 # to avoid numerical errors
num = np.sum(np.square(references), axis=(1, 2))
den = np.sum(np.square(references - estimates), axis=(1, 2))
num += eps
den += eps
return 10 * np.log10(num / den)
def si_sdr(reference: np.ndarray, estimate: np.ndarray) -> float:
"""
Compute Scale-Invariant Signal-to-Distortion Ratio (SI-SDR) for one or more audio tracks.
SI-SDR is a variant of the SDR metric that is invariant to the scaling of the estimate relative to the reference.
It is calculated by scaling the estimate to match the reference signal and then computing the SDR.
Parameters:
----------
reference : np.ndarray
A 3D numpy array of shape (num_sources, num_channels, num_samples), where num_sources is the number of sources,
num_channels is the number of channels (e.g., 1 for mono, 2 for stereo), and num_samples is the length of the audio signal.
estimate : np.ndarray
A 3D numpy array of shape (num_sources, num_channels, num_samples) representing the estimated sources.
Returns:
-------
float
The SI-SDR value for the source. It is a scalar representing the Signal-to-Distortion Ratio in decibels (dB).
"""
eps = 1e-8 # To avoid numerical errors
scale = np.sum(estimate * reference + eps, axis=(0, 1)) / np.sum(reference ** 2 + eps, axis=(0, 1))
scale = np.expand_dims(scale, axis=(0, 1)) # Reshape to [num_sources, 1]
reference = reference * scale
si_sdr = np.mean(10 * np.log10(
np.sum(reference ** 2, axis=(0, 1)) / (np.sum((reference - estimate) ** 2, axis=(0, 1)) + eps) + eps))
return si_sdr
def L1Freq_metric(
reference: np.ndarray,
estimate: np.ndarray,
fft_size: int = 2048,
hop_size: int = 1024,
device: str = 'cpu'
) -> float:
"""
Compute the L1 Frequency Metric between the reference and estimated audio signals.
This metric compares the magnitude spectrograms of the reference and estimated audio signals
using the Short-Time Fourier Transform (STFT) and calculates the L1 loss between them. The result
is scaled to the range [0, 100] where a higher value indicates better performance.
Parameters:
----------
reference : np.ndarray
A 2D numpy array of shape (num_channels, num_samples) representing the reference (ground truth) audio signal.
estimate : np.ndarray
A 2D numpy array of shape (num_channels, num_samples) representing the estimated (predicted) audio signal.
fft_size : int, optional
The size of the FFT (Short-Time Fourier Transform). Default is 2048.
hop_size : int, optional
The hop size between STFT frames. Default is 1024.
device : str, optional
The device to run the computation on ('cpu' or 'cuda'). Default is 'cpu'.
Returns:
-------
float
The L1 Frequency Metric in the range [0, 100], where higher values indicate better performance.
"""
reference = torch.from_numpy(reference).to(device)
estimate = torch.from_numpy(estimate).to(device)
reference_stft = torch.stft(reference, fft_size, hop_size, return_complex=True)
estimated_stft = torch.stft(estimate, fft_size, hop_size, return_complex=True)
reference_mag = torch.abs(reference_stft)
estimate_mag = torch.abs(estimated_stft)
loss = 10 * F.l1_loss(estimate_mag, reference_mag)
ret = 100 / (1. + float(loss.cpu().numpy()))
return ret
def LogWMSE_metric(
reference: np.ndarray,
estimate: np.ndarray,
mixture: np.ndarray,
device: str = 'cpu',
) -> float:
"""
Calculate the Log-WMSE (Logarithmic Weighted Mean Squared Error) between the reference, estimate, and mixture signals.
This metric evaluates the quality of the estimated signal compared to the reference signal in the
context of audio source separation. The result is given in logarithmic scale, which helps in evaluating
signals with large amplitude differences.
Parameters:
----------
reference : np.ndarray
The ground truth audio signal of shape (channels, time), where channels is the number of audio channels
(e.g., 1 for mono, 2 for stereo) and time is the length of the audio in samples.
estimate : np.ndarray
The estimated audio signal of shape (channels, time).
mixture : np.ndarray
The mixed audio signal of shape (channels, time).
device : str, optional
The device to run the computation on, either 'cpu' or 'cuda'. Default is 'cpu'.
Returns:
-------
float
The Log-WMSE value, which quantifies the difference between the reference and estimated signal on a logarithmic scale.
"""
from torch_log_wmse import LogWMSE
log_wmse = LogWMSE(
audio_length=reference.shape[-1] / 44100, # audio length in seconds
sample_rate=44100, # sample rate of 44100 Hz
return_as_loss=False, # return as loss (False means return as metric)
bypass_filter=False, # bypass frequency filtering (False means apply filter)
)
reference = torch.from_numpy(reference).unsqueeze(0).unsqueeze(0).to(device)
estimate = torch.from_numpy(estimate).unsqueeze(0).unsqueeze(0).to(device)
mixture = torch.from_numpy(mixture).unsqueeze(0).to(device)
res = log_wmse(mixture, reference, estimate)
return float(res.cpu().numpy())
def AuraSTFT_metric(
reference: np.ndarray,
estimate: np.ndarray,
device: str = 'cpu',
) -> float:
"""
Calculate the AuraSTFT metric, which evaluates the spectral difference between the reference and estimated
audio signals using Short-Time Fourier Transform (STFT) loss.
The AuraSTFT metric computes the STFT loss in both logarithmic and linear magnitudes, and it is commonly used
to assess the quality of audio separation tasks. The result is returned as a value scaled to the range [0, 100].
Parameters:
----------
reference : np.ndarray
The ground truth audio signal of shape (channels, time), where channels is the number of audio channels
(e.g., 1 for mono, 2 for stereo) and time is the length of the audio in samples.
estimate : np.ndarray
The estimated audio signal of shape (channels, time).
device : str, optional
The device to run the computation on, either 'cpu' or 'cuda'. Default is 'cpu'.
Returns:
-------
float
The AuraSTFT metric value, scaled to the range [0, 100], which quantifies the difference between
the reference and estimated signal in the spectral domain.
"""
from auraloss.freq import STFTLoss
stft_loss = STFTLoss(
w_log_mag=1.0, # weight for log magnitude
w_lin_mag=0.0, # weight for linear magnitude
w_sc=1.0, # weight for spectral centroid
device=device,
)
reference = torch.from_numpy(reference).unsqueeze(0).to(device)
estimate = torch.from_numpy(estimate).unsqueeze(0).to(device)
res = 100 / (1. + 10 * stft_loss(reference, estimate))
return float(res.cpu().numpy())
def AuraMRSTFT_metric(
reference: np.ndarray,
estimate: np.ndarray,
device: str = 'cpu',
) -> float:
"""
Calculate the AuraMRSTFT metric, which evaluates the spectral difference between the reference and estimated
audio signals using Multi-Resolution Short-Time Fourier Transform (STFT) loss.
The AuraMRSTFT metric uses multi-resolution STFT analysis, which allows better representation of both
low- and high-frequency components in the audio signals. The result is returned as a value scaled to the range [0, 100].
Parameters:
----------
reference : np.ndarray
The ground truth audio signal of shape (channels, time), where channels is the number of audio channels
(e.g., 1 for mono, 2 for stereo) and time is the length of the audio in samples.
estimate : np.ndarray
The estimated audio signal of shape (channels, time).
device : str, optional
The device to run the computation on, either 'cpu' or 'cuda'. Default is 'cpu'.
Returns:
-------
float
The AuraMRSTFT metric value, scaled to the range [0, 100], which quantifies the difference between
the reference and estimated signal in the multi-resolution spectral domain.
"""
from auraloss.freq import MultiResolutionSTFTLoss
mrstft_loss = MultiResolutionSTFTLoss(
fft_sizes=[1024, 2048, 4096],
hop_sizes=[256, 512, 1024],
win_lengths=[1024, 2048, 4096],
scale="mel", # mel scale for frequency resolution
n_bins=128, # number of bins for mel scale
sample_rate=44100,
perceptual_weighting=True, # apply perceptual weighting
device=device
)
reference = torch.from_numpy(reference).unsqueeze(0).float().to(device)
estimate = torch.from_numpy(estimate).unsqueeze(0).float().to(device)
res = 100 / (1. + 10 * mrstft_loss(reference, estimate))
return float(res.cpu().numpy())
def bleed_full(
reference: np.ndarray,
estimate: np.ndarray,
sr: int = 44100,
n_fft: int = 4096,
hop_length: int = 1024,
n_mels: int = 512,
device: str = 'cpu',
) -> Tuple[float, float]:
"""
Calculate the 'bleed' and 'fullness' metrics between a reference and an estimated audio signal.
The 'bleed' metric measures how much the estimated signal bleeds into the reference signal,
while the 'fullness' metric measures how much the estimated signal retains its distinctiveness
in relation to the reference signal, both using mel spectrograms and decibel scaling.
Parameters:
----------
reference : np.ndarray
The reference audio signal, shape (channels, time), where channels is the number of audio channels
(e.g., 1 for mono, 2 for stereo) and time is the length of the audio in samples.
estimate : np.ndarray
The estimated audio signal, shape (channels, time).
sr : int, optional
The sample rate of the audio signals. Default is 44100 Hz.
n_fft : int, optional
The FFT size used to compute the STFT. Default is 4096.
hop_length : int, optional
The hop length for STFT computation. Default is 1024.
n_mels : int, optional
The number of mel frequency bins. Default is 512.
device : str, optional
The device for computation, either 'cpu' or 'cuda'. Default is 'cpu'.
Returns:
-------
tuple
A tuple containing two values:
- `bleedless` (float): A score indicating how much 'bleeding' the estimated signal has (higher is better).
- `fullness` (float): A score indicating how 'full' the estimated signal is (higher is better).
"""
from torchaudio.transforms import AmplitudeToDB
reference = torch.from_numpy(reference).float().to(device)
estimate = torch.from_numpy(estimate).float().to(device)
window = torch.hann_window(n_fft).to(device)
# Compute STFTs with the Hann window
D1 = torch.abs(torch.stft(reference, n_fft=n_fft, hop_length=hop_length, window=window, return_complex=True,
pad_mode="constant"))
D2 = torch.abs(torch.stft(estimate, n_fft=n_fft, hop_length=hop_length, window=window, return_complex=True,
pad_mode="constant"))
mel_basis = librosa.filters.mel(sr=sr, n_fft=n_fft, n_mels=n_mels)
mel_filter_bank = torch.from_numpy(mel_basis).to(device)
S1_mel = torch.matmul(mel_filter_bank, D1)
S2_mel = torch.matmul(mel_filter_bank, D2)
S1_db = AmplitudeToDB(stype="magnitude", top_db=80)(S1_mel)
S2_db = AmplitudeToDB(stype="magnitude", top_db=80)(S2_mel)
diff = S2_db - S1_db
positive_diff = diff[diff > 0]
negative_diff = diff[diff < 0]
average_positive = torch.mean(positive_diff) if positive_diff.numel() > 0 else torch.tensor(0.0).to(device)
average_negative = torch.mean(negative_diff) if negative_diff.numel() > 0 else torch.tensor(0.0).to(device)
bleedless = 100 * 1 / (average_positive + 1)
fullness = 100 * 1 / (-average_negative + 1)
return bleedless.cpu().numpy(), fullness.cpu().numpy()
def get_metrics(
metrics: List[str],
reference: np.ndarray,
estimate: np.ndarray,
mix: np.ndarray,
device: str = 'cpu',
) -> Dict[str, float]:
"""
Calculate a list of metrics to evaluate the performance of audio source separation models.
The function computes the specified metrics based on the reference, estimate, and mixture.
Parameters:
----------
metrics : List[str]
A list of metric names to compute (e.g., ['sdr', 'si_sdr', 'l1_freq']).
reference : np.ndarray
The reference audio (true signal) with shape (channels, length).
estimate : np.ndarray
The estimated audio (predicted signal) with shape (channels, length).
mix : np.ndarray
The mixed audio signal with shape (channels, length).
device : str, optional, default='cpu'
The device ('cpu' or 'cuda') to perform the calculations on.
Returns:
-------
Dict[str, float]
A dictionary containing the computed metric values.
"""
result = dict()
# Adjust the length to be the same across all inputs
min_length = min(reference.shape[1], estimate.shape[1])
reference = reference[..., :min_length]
estimate = estimate[..., :min_length]
mix = mix[..., :min_length]
if 'sdr' in metrics:
references = np.expand_dims(reference, axis=0)
estimates = np.expand_dims(estimate, axis=0)
result['sdr'] = float(sdr(references, estimates))
if 'si_sdr' in metrics:
result['si_sdr'] = float(si_sdr(reference, estimate))
if 'l1_freq' in metrics:
result['l1_freq'] = L1Freq_metric(reference, estimate, device=device)
if 'log_wmse' in metrics:
result['log_wmse'] = LogWMSE_metric(reference, estimate, mix, device)
if 'aura_stft' in metrics:
result['aura_stft'] = AuraSTFT_metric(reference, estimate, device)
if 'aura_mrstft' in metrics:
result['aura_mrstft'] = AuraMRSTFT_metric(reference, estimate, device)
if 'bleedless' in metrics or 'fullness' in metrics:
bleedless, fullness = bleed_full(reference, estimate, device=device)
if 'bleedless' in metrics:
result['bleedless'] = float(bleedless)
if 'fullness' in metrics:
result['fullness'] = float(fullness)
return result
@@ -0,0 +1,777 @@
# coding: utf-8
__author__ = 'Roman Solovyev (ZFTurbo): https://github.com/ZFTurbo/'
import argparse
import numpy as np
import torch
import torch.nn as nn
from ml_collections import ConfigDict
from torch.optim import Adam, AdamW, SGD, RAdam, RMSprop
from tqdm.auto import tqdm
from typing import Dict, List, Tuple, Any, Union, Optional
import loralib as lora
from .muon import SingleDeviceMuonWithAuxAdam
import torch.distributed as dist
def demix(
config: ConfigDict,
model: torch.nn.Module,
mix: torch.Tensor,
device: torch.device,
model_type: str,
pbar: bool = False
) -> Union[Dict[str, np.ndarray], np.ndarray]:
"""
Perform audio source separation with a given model.
Supports both Demucs-specific and generic processing modes, including
overlapping chunk-based inference with optional progress bar display.
Handles padding, fading, and batching to reduce artifacts during separation.
Args:
config (ConfigDict): Configuration object with audio and inference
parameters (chunk size, overlap, batch size, etc.).
model (torch.nn.Module): Source separation model for inference.
mix (torch.Tensor): Input audio tensor of shape (channels, time).
device (torch.device): Device on which to run inference (CPU or CUDA).
model_type (str): Type of model (e.g., 'htdemucs', 'mdx23c') that
determines processing mode.
pbar (bool, optional): If True, show a progress bar during chunk
processing. Defaults to False.
Returns:
Union[Dict[str, np.ndarray], np.ndarray]:
- Dictionary mapping instrument names to separated waveforms if
multiple instruments are predicted.
- NumPy array of separated audio if only a single instrument is
present (Demucs mode).
"""
should_print = not dist.is_initialized() or dist.get_rank() == 0
mix = torch.tensor(mix, dtype=torch.float32)
if model_type == 'htdemucs':
mode = 'demucs'
else:
mode = 'generic'
# Define processing parameters based on the mode
if mode == 'demucs':
chunk_size = config.training.samplerate * config.training.segment
num_instruments = len(config.training.instruments)
num_overlap = config.inference.num_overlap
step = chunk_size // num_overlap
else:
if 'chunk_size' in config.inference:
chunk_size = config.inference.chunk_size
else:
chunk_size = config.audio.chunk_size
num_instruments = len(prefer_target_instrument(config))
num_overlap = config.inference.num_overlap
fade_size = chunk_size // 10
step = chunk_size // num_overlap
border = chunk_size - step
length_init = mix.shape[-1]
windowing_array = _getWindowingArray(chunk_size, fade_size)
# Add padding for generic mode to handle edge artifacts
if length_init > 2 * border and border > 0:
mix = nn.functional.pad(mix, (border, border), mode="reflect")
batch_size = config.inference.batch_size
use_amp = getattr(config.training, 'use_amp', True)
with torch.cuda.amp.autocast(enabled=use_amp):
with torch.inference_mode():
# Initialize result and counter tensors
req_shape = (num_instruments,) + mix.shape
result = torch.zeros(req_shape, dtype=torch.float32)
counter = torch.zeros(req_shape, dtype=torch.float32)
i = 0
batch_data = []
batch_locations = []
if pbar and should_print:
progress_bar = tqdm(
total=mix.shape[1], desc="Processing audio chunks", leave=False
)
else:
progress_bar = None
while i < mix.shape[1]:
# Extract chunk and apply padding if necessary
part = mix[:, i:i + chunk_size].to(device)
chunk_len = part.shape[-1]
if mode == "generic" and chunk_len > chunk_size // 2:
pad_mode = "reflect"
else:
pad_mode = "constant"
part = nn.functional.pad(part, (0, chunk_size - chunk_len), mode=pad_mode, value=0)
batch_data.append(part)
batch_locations.append((i, chunk_len))
i += step
# Process batch if it's full or the end is reached
if len(batch_data) >= batch_size or i >= mix.shape[1]:
arr = torch.stack(batch_data, dim=0)
x = model(arr)
if mode == "generic":
window = windowing_array.clone() # using clone() fixes the clicks at chunk edges when using batch_size=1
if i - step == 0: # First audio chunk, no fadein
window[:fade_size] = 1
elif i >= mix.shape[1]: # Last audio chunk, no fadeout
window[-fade_size:] = 1
for j, (start, seg_len) in enumerate(batch_locations):
if mode == "generic":
result[..., start:start + seg_len] += x[j, ..., :seg_len].cpu() * window[..., :seg_len]
counter[..., start:start + seg_len] += window[..., :seg_len]
else:
result[..., start:start + seg_len] += x[j, ..., :seg_len].cpu()
counter[..., start:start + seg_len] += 1.0
batch_data.clear()
batch_locations.clear()
if progress_bar:
progress_bar.update(step)
if progress_bar:
progress_bar.close()
"""
# mix: B, 2, T
# req_shape = (num_instruments,) + mix.shape
req_shape = (num_instruments,) + mix.shape
result = torch.zeros(req_shape, dtype=torch.float32)
counter = torch.zeros(req_shape, dtype=torch.float32)
# prev_i = 0
i = 0
batch_data = []
batch_locations = []
while i < mix.shape[-1]:
part = mix[:, :, i:i + chunk_size].to(device)
chunk_len = part.shape[-1]
if mode == "generic" and chunk_len > chunk_size // 2:
pad_mode = "reflect"
else:
pad_mode = "constant"
part = nn.functional.pad(part, (0, chunk_size - chunk_len), mode=pad_mode, value=0)
# batch_locations.append((i, chunk_len))
# prev_i = i
batch_location = i, i + chunk_len
i += step
# print(part.shape)
x = model(part)
x = x.transpose(0, 1)
# print(x.shape)
if mode == "generic":
window = windowing_array.clone() # using clone() fixes the clicks at chunk edges when using batch_size=1
if i - step == 0: # First audio chunk, no fadein
window[:fade_size] = 1
elif i >= mix.shape[1]: # Last audio chunk, no fadeout
window[-fade_size:] = 1
# for j, (start, seg_len) in enumerate(batch_locations):
# l = chunk_len if chunk_len < chunk_size else chunk_size
# print(l, x.shape, result.shape, counter.shape, window.shape)
# print(result[..., batch_location[0]: batch_location[1]].shape, x[..., :chunk_len].cpu().shape, window[..., :chunk_len].shape)
if mode == "generic":
result[..., batch_location[0]: batch_location[1]] += x[..., :chunk_len].cpu() * window[..., :chunk_len]
counter[..., batch_location[0]: batch_location[1]] += window[..., :chunk_len]
else:
result[..., batch_location[0]: batch_location[1]] += x[..., :chunk_len].cpu()
counter[..., batch_location[0]: batch_location[1]] += 1.0
batch_data.clear()
batch_locations.clear()
"""
# Compute final estimated sources
estimated_sources = result / counter
estimated_sources = estimated_sources.cpu().numpy()
np.nan_to_num(estimated_sources, copy=False, nan=0.0)
# Remove padding for generic mode
if mode == "generic":
if length_init > 2 * border and border > 0:
estimated_sources = estimated_sources[..., border:-border]
# Return the result as a dictionary or a single array
if mode == "demucs":
instruments = config.training.instruments
else:
instruments = prefer_target_instrument(config)
ret_data = {k: v for k, v in zip(instruments, estimated_sources)}
if mode == "demucs" and num_instruments <= 1:
return estimated_sources
else:
return ret_data
def initialize_model_and_device(model: torch.nn.Module, device_ids: List[int]) -> Tuple[Union[torch.device, str], torch.nn.Module]:
"""
Move a model to the correct computation device and wrap with DataParallel if needed.
Selects GPU(s) if CUDA is available; otherwise defaults to CPU. If multiple
GPU IDs are provided, wraps the model with `nn.DataParallel` for multi-GPU
execution.
Args:
model (torch.nn.Module): PyTorch model to be initialized.
device_ids (List[int]): List of GPU device IDs to use. If length > 1,
the model will be wrapped with DataParallel.
Returns:
Tuple[Union[torch.device, str], torch.nn.Module]: A tuple containing:
- The computation device (`torch.device` or "cpu").
- The model moved to that device (wrapped in DataParallel if applicable).
"""
if torch.cuda.is_available():
if len(device_ids) <= 1:
device = torch.device(f'cuda:{device_ids[0]}')
model = model.to(device)
else:
device = torch.device(f'cuda:{device_ids[0]}')
model = nn.DataParallel(model, device_ids=device_ids).to(device)
else:
device = 'cpu'
model = model.to(device)
print("CUDA is not available. Running on CPU.")
return device, model
def get_optimizer(config: ConfigDict, model: torch.nn.Module) -> torch.optim.Optimizer:
"""
Create and configure an optimizer for training.
Selects the optimizer type based on `config.training.optimizer` and applies
the corresponding parameters, including support for advanced optimizers
such as Muon, Prodigy, and 8-bit AdamW. Handles parameter group separation
for specialized optimizers (e.g., Muon vs. Adam parameters).
Args:
config (ConfigDict): Training configuration containing optimizer type,
learning rate, and optional optimizer-specific parameters.
model (torch.nn.Module): Model whose parameters will be optimized.
Returns:
torch.optim.Optimizer: Initialized optimizer ready for training.
Raises:
ValueError: If required optimizer configuration is missing (e.g., for Muon).
SystemExit: If an unknown optimizer name is encountered.
"""
should_print = not dist.is_initialized() or dist.get_rank() == 0
optim_params = dict()
if 'optimizer' in config:
optim_params = dict(config['optimizer'])
if config.training.optimizer != 'muon' and should_print:
print(f'Optimizer params from config:\n{optim_params}')
name_optimizer = getattr(config.training, 'optimizer',
'No optimizer in config')
if name_optimizer == 'adam':
optimizer = Adam(model.parameters(), lr=config.training.lr, **optim_params)
elif name_optimizer == 'adamw':
optimizer = AdamW(model.parameters(), lr=config.training.lr, **optim_params)
elif name_optimizer == 'radam':
optimizer = RAdam(model.parameters(), lr=config.training.lr, **optim_params)
elif name_optimizer == 'rmsprop':
optimizer = RMSprop(model.parameters(), lr=config.training.lr, **optim_params)
elif name_optimizer == 'prodigy':
from prodigyopt import Prodigy
# you can choose weight decay value based on your problem, 0 by default
# We recommend using lr=1.0 (default) for all networks.
optimizer = Prodigy(model.parameters(), lr=config.training.lr, **optim_params)
elif name_optimizer == 'adamw8bit':
import bitsandbytes as bnb
optimizer = bnb.optim.AdamW8bit(model.parameters(), lr=config.training.lr, **optim_params)
elif name_optimizer == 'muon':
if should_print:
print("Using Muon optimizer (Single-Device) with AdamW for auxiliary parameters.")
muon_params = [p for p in model.parameters() if p.ndim >= 2]
adam_params = [p for p in model.parameters() if p.ndim < 2]
if not hasattr(config, 'optimizer') or 'muon_group' not in config.optimizer or 'adam_group' not in config.optimizer:
raise ValueError("For the 'muon' optimizer, the config must have an 'optimizer' section "
"with 'muon_group' and 'adam_group' dictionaries.")
muon_group_config = dict(config.optimizer.muon_group)
adam_group_config = dict(config.optimizer.adam_group)
if should_print:
print(f"Muon group params: {muon_group_config}")
print(f"Adam group params: {adam_group_config}")
param_groups = [
dict(params=muon_params, use_muon=True, **muon_group_config),
dict(params=adam_params, use_muon=False, **adam_group_config),
]
optimizer = SingleDeviceMuonWithAuxAdam(param_groups)
elif name_optimizer == 'sgd':
if should_print:
print('Use SGD optimizer')
optimizer = SGD(model.parameters(), lr=config.training.lr, **optim_params)
else:
if should_print:
print(f'Unknown optimizer: {name_optimizer}')
exit()
return optimizer
def normalize_batch(x: torch.Tensor, y: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Apply mean-variance normalization to a pair of tensors.
Computes the mean and standard deviation from `x` and normalizes both `x`
and `y` using those statistics. This ensures the two tensors are scaled
consistently.
Args:
x (torch.Tensor): Input tensor used to compute normalization statistics.
y (torch.Tensor): Input tensor normalized using the same statistics as `x`.
Returns:
Tuple[torch.Tensor, torch.Tensor]: Normalized tensors `(x, y)`.
"""
mean = x.mean()
std = x.std()
if std != 0:
x = (x - mean) / std
y = (y - mean) / std
return x, y
def apply_tta(
config,
model: torch.nn.Module,
mix: torch.Tensor,
waveforms_orig: Dict[str, torch.Tensor],
device: torch.device,
model_type: str
) -> Dict[str, torch.Tensor]:
"""
Enhance source separation results using Test-Time Augmentation (TTA).
Applies augmentations such as channel reversal and polarity inversion to
the input mixture, reprocesses with the model, and combines the results
with the original predictions by averaging.
Args:
config: Configuration object with model and inference parameters.
model (torch.nn.Module): Trained source separation model.
mix (torch.Tensor): Input mixture tensor of shape (channels, time).
waveforms_orig (Dict[str, torch.Tensor]): Dictionary of separated
sources before augmentation.
device (torch.device): Computation device (CPU or CUDA).
model_type (str): Model type identifier used for demixing.
Returns:
Dict[str, torch.Tensor]: Dictionary of separated sources after applying TTA.
"""
# Create augmentations: channel inversion and polarity inversion
track_proc_list = [mix[::-1].copy(), -1.0 * mix.copy()]
# Process each augmented mixture
for i, augmented_mix in enumerate(track_proc_list):
waveforms = demix(config, model, augmented_mix, device, model_type=model_type)
for el in waveforms:
if i == 0:
waveforms_orig[el] += waveforms[el][::-1].copy()
else:
waveforms_orig[el] -= waveforms[el]
# Average the results across augmentations
for el in waveforms_orig:
waveforms_orig[el] /= len(track_proc_list) + 1
return waveforms_orig
def _getWindowingArray(window_size: int, fade_size: int) -> torch.Tensor:
"""
Generate a windowing array with a linear fade-in at the beginning and a fade-out at the end.
This function creates a window of size `window_size` where the first `fade_size` elements
linearly increase from 0 to 1 (fade-in) and the last `fade_size` elements linearly decrease
from 1 to 0 (fade-out). The middle part of the window is filled with ones.
Parameters:
----------
window_size : int
The total size of the window.
fade_size : int
The size of the fade-in and fade-out regions.
Returns:
-------
torch.Tensor
A tensor of shape (window_size,) containing the generated windowing array.
Example:
-------
If `window_size=10` and `fade_size=3`, the output will be:
tensor([0.0000, 0.5000, 1.0000, 1.0000, 1.0000, 1.0000, 1.0000, 1.0000, 0.5000, 0.0000])
"""
fadein = torch.linspace(0, 1, fade_size)
fadeout = torch.linspace(1, 0, fade_size)
window = torch.ones(window_size)
window[-fade_size:] = fadeout
window[:fade_size] = fadein
return window
def prefer_target_instrument(config: ConfigDict) -> List[str]:
"""
Return the list of target instruments based on the configuration.
If a specific target instrument is specified in the configuration,
it returns a list with that instrument. Otherwise, it returns the list of instruments.
Parameters:
----------
config : ConfigDict
Configuration object containing the list of instruments or the target instrument.
Returns:
-------
List[str]
A list of target instruments.
"""
if getattr(config.training, 'target_instrument', None):
return [config.training.target_instrument]
else:
return config.training.instruments
def load_not_compatible_weights(model: torch.nn.Module, old_model: dict, verbose: bool = False) -> None:
"""
Load a possibly incompatible state dict into `model` with best-effort matching.
Accepts either a raw state_dict or a checkpoint dict with weights under "state" or "state_dict".
For each param/buffer in `model`: if the name exists and shapes match → copy;
if ndim matches but shapes differ → zero-pad/crop the source to fit the target;
if the name is missing or ndim differs → skip. Optional logging on rank 0 when `verbose=True`.
Args:
model: Target PyTorch module.
old_model: Source weights (state_dict or checkpoint dict).
verbose: Print brief load decisions.
Returns:
None
"""
should_print = verbose and (not dist.is_initialized() or dist.get_rank() == 0)
new_model = model.state_dict()
if 'state' in old_model:
# Fix for htdemucs weights loading
old_model = old_model['state']
if 'state_dict' in old_model:
# Fix for apollo weights loading
old_model = old_model['state_dict']
if 'model_state_dict' in old_model:
# Fix for full_check_point
old_model = old_model['model_state_dict']
for el in new_model:
if el in old_model:
if should_print:
print(f'Match found for {el}!')
if new_model[el].shape == old_model[el].shape:
if should_print:
print('Action: Just copy weights!')
new_model[el] = old_model[el]
else:
if len(new_model[el].shape) != len(old_model[el].shape) and should_print:
print('Action: Different dimension! Too lazy to write the code... Skip it')
else:
if should_print:
print(f'Shape is different: {tuple(new_model[el].shape)} != {tuple(old_model[el].shape)}')
ln = len(new_model[el].shape)
max_shape = []
slices_old = []
slices_new = []
for i in range(ln):
max_shape.append(max(new_model[el].shape[i], old_model[el].shape[i]))
slices_old.append(slice(0, old_model[el].shape[i]))
slices_new.append(slice(0, new_model[el].shape[i]))
# print(max_shape)
# print(slices_old, slices_new)
slices_old = tuple(slices_old)
slices_new = tuple(slices_new)
max_matrix = np.zeros(max_shape, dtype=np.float32)
for i in range(ln):
max_matrix[slices_old] = old_model[el].cpu().numpy()
max_matrix = torch.from_numpy(max_matrix)
new_model[el] = max_matrix[slices_new]
else:
if should_print:
print(f'Match not found for {el}!')
model.load_state_dict(
new_model
)
def load_lora_weights(model: torch.nn.Module, lora_path: str, device: str = 'cpu') -> None:
"""
Load LoRA weights into a model.
This function updates the given model with LoRA-specific weights from the specified checkpoint file.
It does not require the checkpoint to match the model's full state dictionary, as only LoRA layers are updated.
Parameters:
----------
model : Module
The PyTorch model into which the LoRA weights will be loaded.
lora_path : str
Path to the LoRA checkpoint file.
device : str, optional
The device to load the weights onto, by default 'cpu'. Common values are 'cpu' or 'cuda'.
Returns:
-------
None
The model is updated in place.
"""
lora_state_dict = torch.load(lora_path, map_location=device)
model.load_state_dict(lora_state_dict, strict=False)
def load_start_checkpoint(args: argparse.Namespace,
model: torch.nn.Module,
old_model: None,
type_: str = 'train') -> None:
"""
Load an initial checkpoint into `model`.
For `type_ == "train"`, performs a tolerant load using `old_model` (a state dict or a
checkpoint dict) via `load_not_compatible_weights`, allowing partial shape mismatches.
For other modes, loads a strict state dict from `args.start_check_point`, with special
handling for HTDemucs/Apollo checkpoints (keys under "state"/"state_dict"). If
`args.lora_checkpoint` is set, LoRA weights are applied after the base load.
Args:
args: Namespace with at least `start_check_point`, `model_type`, and optionally `lora_checkpoint`.
model: Target PyTorch module to receive weights.
old_model: Source weights for tolerant loading in train mode (state dict or checkpoint dict).
type_: Loading strategy; "train" uses tolerant loading, otherwise strict loading from path.
Returns:
None
"""
should_print = not dist.is_initialized() or dist.get_rank() == 0
if should_print:
print(f'Start from checkpoint: {args.start_check_point}')
if type_ in ['train']:
if 1:
load_not_compatible_weights(model, old_model, verbose=False)
else:
model.load_state_dict(torch.load(args.start_check_point))
else:
device='cpu'
if args.model_type in ['htdemucs', 'apollo']:
state_dict = torch.load(args.start_check_point, map_location=device, weights_only=False)
# Fix for htdemucs pretrained models
if 'state' in state_dict:
state_dict = state_dict['state']
# Fix for apollo pretrained models
if 'state_dict' in state_dict:
state_dict = state_dict['state_dict']
else:
state_dict = torch.load(args.start_check_point, map_location=device, weights_only=True)
model.load_state_dict(state_dict)
if args.lora_checkpoint:
if should_print:
print(f"Loading LoRA weights from: {args.lora_checkpoint}")
load_lora_weights(model, args.lora_checkpoint)
def bind_lora_to_model(config: Dict[str, Any], model: nn.Module) -> nn.Module:
"""
Replaces specific layers in the model with LoRA-extended versions.
Parameters:
----------
config : Dict[str, Any]
Configuration containing parameters for LoRA. It should include a 'lora' key with parameters for `MergedLinear`.
model : nn.Module
The original model in which the layers will be replaced.
Returns:
-------
nn.Module
The modified model with the replaced layers.
"""
if 'lora' not in config:
raise ValueError("Configuration must contain the 'lora' key with parameters for LoRA.")
replaced_layers = 0 # Counter for replaced layers
should_print = not dist.is_initialized() or dist.get_rank() == 0
for name, module in model.named_modules():
hierarchy = name.split('.')
layer_name = hierarchy[-1]
# Check if this is the target layer to replace (and layer_name == 'to_qkv')
if isinstance(module, nn.Linear):
try:
# Get the parent module
parent_module = model
for submodule_name in hierarchy[:-1]:
parent_module = getattr(parent_module, submodule_name)
# Replace the module with LoRA-enabled layer
setattr(
parent_module,
layer_name,
lora.MergedLinear(
in_features=module.in_features,
out_features=module.out_features,
bias=module.bias is not None,
**config['lora']
)
)
replaced_layers += 1 # Increment the counter
except Exception as e:
if should_print:
print(f"Error replacing layer {name}: {e}")
if replaced_layers == 0 and should_print:
print("Warning: No layers were replaced. Check the model structure and configuration.")
elif should_print:
print(f"Number of layers replaced with LoRA: {replaced_layers}")
return model
def save_weights(
store_path: str,
model: nn.Module,
device_ids: List[int],
optimizer: torch.optim.Optimizer,
epoch: int,
all_time_all_metrics,
best_metric: float,
scheduler: Optional[torch.optim.lr_scheduler.ReduceLROnPlateau] = None,
train_lora: bool = False
) -> None:
"""
Save a training checkpoint containing model weights, optimizer/scheduler states, and metadata.
Behavior:
- In Distributed Data Parallel (DDP), only rank 0 writes the file to avoid conflicts.
- If `train_lora` is True, saves only LoRA adapter weights (`lora_state_dict`); otherwise saves the full model.
- Uses `model.module.state_dict()` when the model is wrapped by DDP/DataParallel.
- Stores `epoch` and `best_metric` alongside optimizer/scheduler states.
Args:
store_path: Destination file path for the checkpoint (will be overwritten).
model: The model whose weights are being saved (may be wrapped by DDP/DataParallel).
device_ids: List of GPU device IDs used during training (used to detect DP wrapping in non-DDP runs).
optimizer: Optimizer whose state will be saved.
epoch: Current training epoch to record in the checkpoint.
all_time_all_metrics:
best_metric: Best validation metric achieved so far.
scheduler: Optional learning rate scheduler; its state is saved if provided.
train_lora: If True, save only LoRA adapter weights instead of the full model.
Returns:
None
"""
checkpoint: Dict[str, Any] = {
"epoch": epoch,
"optimizer_name": optimizer.__class__.__name__,
"optimizer_state_dict": optimizer.state_dict(),
"scheduler_state_dict": scheduler.state_dict() if scheduler else None,
"best_metric": best_metric,
"all_metrics": all_time_all_metrics
}
# Save model weights
if train_lora:
checkpoint["model_state_dict"] = lora.lora_state_dict(model)
else:
if dist.is_initialized():
# In DDP, use .module
checkpoint["model_state_dict"] = model.module.state_dict()
else:
checkpoint["model_state_dict"] = (
model.state_dict() if len(device_ids) <= 1 else model.module.state_dict()
)
# Save only on rank 0 (or if not using DDP)
if not dist.is_initialized() or dist.get_rank() == 0:
torch.save(checkpoint, store_path)
def save_last_weights(
args: argparse.Namespace,
model: nn.Module,
device_ids: List[int],
optimizer: torch.optim.Optimizer,
epoch: int,
all_time_all_metrics,
best_metric: float,
scheduler: Optional[torch.optim.lr_scheduler.ReduceLROnPlateau] = None,
) -> None:
"""
Save the latest training checkpoint for continuation or recovery.
The checkpoint is always written to:
{args.results_path}/last_{args.model_type}.ckpt
This wraps `save_weights` and ensures the latest model/optimizer/scheduler
states are recorded, along with the current epoch and best metric. In DDP,
only rank 0 performs the save. Supports both standard and LoRA training.
Args:
all_time_all_metrics:
args: Training arguments. Must define `results_path`, `model_type`,
and `train_lora`.
model: Model instance (may be wrapped by DDP/DataParallel).
device_ids: List of GPU IDs used for training.
optimizer: Optimizer whose state will be saved.
epoch: Current training epoch.
best_metric: Current best validation metric.
scheduler: Optional learning rate scheduler to save state for.
Returns:
None
"""
store_path = f"{args.results_path}/last_{args.model_type}.ckpt"
save_weights(
store_path,
model,
device_ids,
optimizer,
epoch,
all_time_all_metrics,
best_metric,
scheduler,
args.train_lora,
)
@@ -0,0 +1,286 @@
import torch
import torch.distributed as dist
def zeropower_via_newtonschulz5(G, steps: int):
"""
Newton-Schulz iteration to compute the zeroth power / orthogonalization of G. We opt to use a
quintic iteration whose coefficients are selected to maximize the slope at zero. For the purpose
of minimizing steps, it turns out to be empirically effective to keep increasing the slope at
zero even beyond the point where the iteration no longer converges all the way to one everywhere
on the interval. This iteration therefore does not produce UV^T but rather something like US'V^T
where S' is diagonal with S_{ii}' ~ Uniform(0.5, 1.5), which turns out not to hurt model
performance at all relative to UV^T, where USV^T = G is the SVD.
"""
assert G.ndim >= 2 # batched Muon implementation by @scottjmaddox, and put into practice in the record by @YouJiacheng
a, b, c = (3.4445, -4.7750, 2.0315)
X = G.bfloat16()
if G.size(-2) > G.size(-1):
X = X.mT
# Ensure spectral norm is at most 1
X = X / (X.norm(dim=(-2, -1), keepdim=True) + 1e-7)
# Perform the NS iterations
for _ in range(steps):
A = X @ X.mT
B = b * A + c * A @ A # quintic computation strategy adapted from suggestion by @jxbz, @leloykun, and @YouJiacheng
X = a * X + B @ X
if G.size(-2) > G.size(-1):
X = X.mT
return X
def muon_update(grad, momentum, beta=0.95, ns_steps=5, nesterov=True):
momentum.lerp_(grad, 1 - beta)
update = grad.lerp_(momentum, beta) if nesterov else momentum
if update.ndim == 4: # for the case of conv filters
update = update.view(len(update), -1)
update = zeropower_via_newtonschulz5(update, steps=ns_steps)
update *= max(1, grad.size(-2) / grad.size(-1))**0.5
return update
class Muon(torch.optim.Optimizer):
"""
Muon - MomentUm Orthogonalized by Newton-schulz
https://kellerjordan.github.io/posts/muon/
Muon internally runs standard SGD-momentum, and then performs an orthogonalization post-
processing step, in which each 2D parameter's update is replaced with the nearest orthogonal
matrix. For efficient orthogonalization we use a Newton-Schulz iteration, which has the
advantage that it can be stably run in bfloat16 on the GPU.
Muon should only be used for hidden weight layers. The input embedding, final output layer,
and any internal gains or biases should be optimized using a standard method such as AdamW.
Hidden convolutional weights can be trained using Muon by viewing them as 2D and then
collapsing their last 3 dimensions.
Arguments:
lr: The learning rate, in units of spectral norm per update.
weight_decay: The AdamW-style weight decay.
momentum: The momentum. A value of 0.95 here is usually fine.
"""
def __init__(self, params, lr=0.02, weight_decay=0, momentum=0.95):
defaults = dict(lr=lr, weight_decay=weight_decay, momentum=momentum)
assert isinstance(params, list) and len(params) >= 1 and isinstance(params[0], torch.nn.Parameter)
params = sorted(params, key=lambda x: x.size(), reverse=True)
super().__init__(params, defaults)
@torch.no_grad()
def step(self, closure=None):
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
for group in self.param_groups:
params = group["params"]
params_pad = params + [torch.empty_like(params[-1])] * (dist.get_world_size() - len(params) % dist.get_world_size())
for base_i in range(len(params))[::dist.get_world_size()]:
if base_i + dist.get_rank() < len(params):
p = params[base_i + dist.get_rank()]
if p.grad is None:
# continue
p.grad = torch.zeros_like(p) # Force synchronization
state = self.state[p]
if len(state) == 0:
state["momentum_buffer"] = torch.zeros_like(p)
update = muon_update(p.grad, state["momentum_buffer"], beta=group["momentum"])
p.mul_(1 - group["lr"] * group["weight_decay"])
p.add_(update.reshape(p.shape), alpha=-group["lr"])
dist.all_gather(params_pad[base_i:base_i + dist.get_world_size()], params_pad[base_i + dist.get_rank()])
return loss
class SingleDeviceMuon(torch.optim.Optimizer):
"""
Muon variant for usage in non-distributed settings.
"""
def __init__(self, params, lr=0.02, weight_decay=0, momentum=0.95):
defaults = dict(lr=lr, weight_decay=weight_decay, momentum=momentum)
super().__init__(params, defaults)
@torch.no_grad()
def step(self, closure=None):
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
for group in self.param_groups:
for p in group["params"]:
if p.grad is None:
# continue
p.grad = torch.zeros_like(p) # Force synchronization
state = self.state[p]
if len(state) == 0:
state["momentum_buffer"] = torch.zeros_like(p)
update = muon_update(p.grad, state["momentum_buffer"], beta=group["momentum"])
p.mul_(1 - group["lr"] * group["weight_decay"])
p.add_(update.reshape(p.shape), alpha=-group["lr"])
return loss
def adam_update(grad, buf1, buf2, step, betas, eps):
buf1.lerp_(grad, 1 - betas[0])
buf2.lerp_(grad.square(), 1 - betas[1])
buf1c = buf1 / (1 - betas[0]**step)
buf2c = buf2 / (1 - betas[1]**step)
return buf1c / (buf2c.sqrt() + eps)
class MuonWithAuxAdam(torch.optim.Optimizer):
"""
Distributed Muon variant that can be used for all parameters in the network, since it runs an
internal AdamW for the parameters that are not compatible with Muon. The user must manually
specify which parameters shall be optimized with Muon and which with Adam by passing in a
list of param_groups with the `use_muon` flag set.
The point of this class is to allow the user to have a single optimizer in their code, rather
than having both a Muon and an Adam which each need to be stepped.
You can see an example usage below:
https://github.com/KellerJordan/modded-nanogpt/blob/master/records/052525_MuonWithAuxAdamExample/b01550f9-03d8-4a9c-86fe-4ab434f1c5e0.txt#L470
```
hidden_matrix_params = [p for n, p in model.blocks.named_parameters() if p.ndim >= 2 and "embed" not in n]
embed_params = [p for n, p in model.named_parameters() if "embed" in n]
scalar_params = [p for p in model.parameters() if p.ndim < 2]
head_params = [model.lm_head.weight]
from muon import MuonWithAuxAdam
adam_groups = [dict(params=head_params, lr=0.22), dict(params=embed_params, lr=0.6), dict(params=scalar_params, lr=0.04)]
adam_groups = [dict(**g, betas=(0.8, 0.95), eps=1e-10, use_muon=False) for g in adam_groups]
muon_group = dict(params=hidden_matrix_params, lr=0.05, momentum=0.95, use_muon=True)
param_groups = [*adam_groups, muon_group]
optimizer = MuonWithAuxAdam(param_groups)
```
"""
def __init__(self, param_groups):
for group in param_groups:
assert "use_muon" in group
if group["use_muon"]:
group["params"] = sorted(group["params"], key=lambda x: x.size(), reverse=True)
# defaults
group["lr"] = group.get("lr", 0.02)
group["momentum"] = group.get("momentum", 0.95)
group["weight_decay"] = group.get("weight_decay", 0)
assert set(group.keys()) == set(["params", "lr", "momentum", "weight_decay", "use_muon"])
else:
# defaults
group["lr"] = group.get("lr", 3e-4)
group["betas"] = group.get("betas", (0.9, 0.95))
group["eps"] = group.get("eps", 1e-10)
group["weight_decay"] = group.get("weight_decay", 0)
assert set(group.keys()) == set(["params", "lr", "betas", "eps", "weight_decay", "use_muon"])
super().__init__(param_groups, dict())
@torch.no_grad()
def step(self, closure=None):
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
for group in self.param_groups:
if group["use_muon"]:
params = group["params"]
params_pad = params + [torch.empty_like(params[-1])] * (dist.get_world_size() - len(params) % dist.get_world_size())
for base_i in range(len(params))[::dist.get_world_size()]:
if base_i + dist.get_rank() < len(params):
p = params[base_i + dist.get_rank()]
if p.grad is None:
# continue
p.grad = torch.zeros_like(p) # Force synchronization
state = self.state[p]
if len(state) == 0:
state["momentum_buffer"] = torch.zeros_like(p)
update = muon_update(p.grad, state["momentum_buffer"], beta=group["momentum"])
p.mul_(1 - group["lr"] * group["weight_decay"])
p.add_(update.reshape(p.shape), alpha=-group["lr"])
dist.all_gather(params_pad[base_i:base_i + dist.get_world_size()], params_pad[base_i + dist.get_rank()])
else:
for p in group["params"]:
if p.grad is None:
# continue
p.grad = torch.zeros_like(p) # Force synchronization
state = self.state[p]
if len(state) == 0:
state["exp_avg"] = torch.zeros_like(p)
state["exp_avg_sq"] = torch.zeros_like(p)
state["step"] = 0
state["step"] += 1
update = adam_update(p.grad, state["exp_avg"], state["exp_avg_sq"],
state["step"], group["betas"], group["eps"])
p.mul_(1 - group["lr"] * group["weight_decay"])
p.add_(update, alpha=-group["lr"])
return loss
class SingleDeviceMuonWithAuxAdam(torch.optim.Optimizer):
"""
Non-distributed variant of MuonWithAuxAdam.
"""
def __init__(self, param_groups):
for group in param_groups:
assert "use_muon" in group
if group["use_muon"]:
# defaults
group["lr"] = group.get("lr", 0.02)
group["momentum"] = group.get("momentum", 0.95)
group["weight_decay"] = group.get("weight_decay", 0)
assert set(group.keys()) == set(["params", "lr", "momentum", "weight_decay", "use_muon"])
else:
# defaults
group["lr"] = group.get("lr", 3e-4)
group["betas"] = group.get("betas", (0.9, 0.95))
group["eps"] = group.get("eps", 1e-10)
group["weight_decay"] = group.get("weight_decay", 0)
assert set(group.keys()) == set(["params", "lr", "betas", "eps", "weight_decay", "use_muon"])
super().__init__(param_groups, dict())
@torch.no_grad()
def step(self, closure=None):
loss = None
if closure is not None:
with torch.enable_grad():
loss = closure()
for group in self.param_groups:
if group["use_muon"]:
for p in group["params"]:
if p.grad is None:
# continue
p.grad = torch.zeros_like(p) # Force synchronization
state = self.state[p]
if len(state) == 0:
state["momentum_buffer"] = torch.zeros_like(p)
update = muon_update(p.grad, state["momentum_buffer"], beta=group["momentum"])
p.mul_(1 - group["lr"] * group["weight_decay"])
p.add_(update.reshape(p.shape), alpha=-group["lr"])
else:
for p in group["params"]:
if p.grad is None:
# continue
p.grad = torch.zeros_like(p) # Force synchronization
state = self.state[p]
if len(state) == 0:
state["exp_avg"] = torch.zeros_like(p)
state["exp_avg_sq"] = torch.zeros_like(p)
state["step"] = 0
state["step"] += 1
update = adam_update(p.grad, state["exp_avg"], state["exp_avg_sq"],
state["step"], group["betas"], group["eps"])
p.mul_(1 - group["lr"] * group["weight_decay"])
p.add_(update, alpha=-group["lr"])
return loss
@@ -0,0 +1,501 @@
import os
import random
import time
import yaml
import wandb
import numpy as np
import torch
import argparse
from typing import Dict, List, Tuple, Union
from omegaconf import OmegaConf
from ml_collections import ConfigDict
import torch.distributed as dist
from torch import nn
def parse_args_train(dict_args: Union[Dict, None]) -> argparse.Namespace:
"""
Parse command-line arguments for training configuration.
This function constructs an argument parser for model, dataset, training, and logging
options, merges overrides from a provided dictionary (if any), and returns the parsed
arguments. If `dict_args` is None, the arguments are parsed from `sys.argv`.
Args:
dict_args (Dict | None): Optional dictionary of argument overrides. Keys should
match the defined CLI options.
Returns:
argparse.Namespace: Parsed arguments namespace containing all configuration
values required for training.
"""
parser = argparse.ArgumentParser()
parser.add_argument("--model_type", type=str, default='mdx23c',
help="One of mdx23c, htdemucs, segm_models, mel_band_roformer, bs_roformer, swin_upernet, bandit")
parser.add_argument("--config_path", type=str, help="path to config file")
parser.add_argument("--start_check_point", type=str, default='', help="Initial checkpoint to start training")
parser.add_argument("--load_optimizer", action='store_true', help="Load optimizer state from checkpoint (if available)")
parser.add_argument("--load_scheduler", action='store_true', help="Load scheduler state from checkpoint (if available)")
parser.add_argument("--load_epoch", action='store_true', help="Load epoch number from checkpoint (if available)")
parser.add_argument("--load_best_metric", action='store_true', help="Load best metric from checkpoint (if available)")
parser.add_argument("--load_all_metrics", action='store_true', help="Load all metrics from checkpoint (if available)")
parser.add_argument("--results_path", type=str,
help="path to folder where results will be stored (weights, metadata)")
parser.add_argument("--data_path", nargs="+", type=str, help="Dataset data paths. You can provide several folders.")
parser.add_argument("--dataset_type", type=int, default=1,
help="Dataset type. Must be one of: 1, 2, 3 or 4. Details here: https://github.com/ZFTurbo/Music-Source-Separation-Training/blob/main/docs/dataset_types.md")
parser.add_argument("--valid_path", nargs="+", type=str,
help="validation data paths. You can provide several folders.")
parser.add_argument("--num_workers", type=int, default=0, help="dataloader num_workers")
parser.add_argument("--pin_memory", action='store_true', help="dataloader pin_memory")
parser.add_argument("--seed", type=int, default=0, help="random seed")
parser.add_argument("--device_ids", nargs='+', type=int, default=[0], help='list of gpu ids')
parser.add_argument("--loss", type=str, nargs='+', choices=['masked_loss', 'mse_loss', 'l1_loss',
'multistft_loss', 'spec_masked_loss', 'spec_rmse_loss', 'log_wmse_loss'],
default=['masked_loss'], help="List of loss functions to use")
parser.add_argument("--masked_loss_coef", type=float, default=1., help="Coef for loss")
parser.add_argument("--mse_loss_coef", type=float, default=1., help="Coef for loss")
parser.add_argument("--l1_loss_coef", type=float, default=1., help="Coef for loss")
parser.add_argument("--log_wmse_loss_coef", type=float, default=1., help="Coef for loss")
parser.add_argument("--multistft_loss_coef", type=float, default=0.001, help="Coef for loss")
parser.add_argument("--spec_masked_loss_coef", type=float, default=1, help="Coef for loss")
parser.add_argument("--spec_rmse_loss_coef", type=float, default=1, help="Coef for loss")
parser.add_argument("--wandb_key", type=str, default='', help='wandb API Key')
parser.add_argument("--wandb_offline", action='store_true', help='local wandb')
parser.add_argument("--pre_valid", action='store_true', help='Run validation before training')
parser.add_argument("--metrics", nargs='+', type=str, default=["sdr"],
choices=['sdr', 'l1_freq', 'si_sdr', 'log_wmse', 'aura_stft', 'aura_mrstft', 'bleedless',
'fullness'], help='List of metrics to use.')
parser.add_argument("--metric_for_scheduler", default="sdr",
choices=['sdr', 'l1_freq', 'si_sdr', 'log_wmse', 'aura_stft', 'aura_mrstft', 'bleedless',
'fullness'], help='Metric which will be used for scheduler.')
parser.add_argument("--train_lora", action='store_true', help="Train with LoRA")
parser.add_argument("--lora_checkpoint", type=str, default='', help="Initial checkpoint to LoRA weights")
parser.add_argument("--each_metrics_in_name", action='store_true',
help="All stems in naming checkpoints")
parser.add_argument("--use_standard_loss", action='store_true',
help="Roformers will use provided loss instead of internal")
parser.add_argument("--save_weights_every_epoch", action='store_true',
help="Weights will be saved every epoch with all metric values")
parser.add_argument("--persistent_workers", action='store_true',
help="dataloader persistent_workers")
parser.add_argument("--prefetch_factor", type=int, default=None,
help="dataloader prefetch_factor")
parser.add_argument("--set_per_process_memory_fraction", action='store_true',
help="using only VRAM, no RAM")
if dict_args is not None:
args = parser.parse_args([])
args_dict = vars(args)
args_dict.update(dict_args)
args = argparse.Namespace(**args_dict)
else:
args = parser.parse_args()
if args.metric_for_scheduler not in args.metrics:
args.metrics += [args.metric_for_scheduler]
get_internal_loss = (args.model_type in ('mel_band_conformer',) or 'roformer' in args.model_type
) and not args.use_standard_loss
if get_internal_loss:
args.loss = [f'{args.model_type}_loss']
return args
def parse_args_valid(dict_args: Union[Dict, None]) -> argparse.Namespace:
"""
Parse command-line arguments for validation configuration.
Builds the CLI for model selection, configuration paths, validation data
locations, output/spectrogram saving options, device/runtime settings, and
evaluation metrics. If `dict_args` is provided, its key–value pairs override
or set the parsed arguments; otherwise arguments are read from `sys.argv`.
Args:
dict_args (Union[Dict, None]): Optional mapping of argument names to values
used to override or supply CLI options programmatically.
Returns:
argparse.Namespace: Parsed arguments namespace containing all validation
configuration values.
"""
parser = argparse.ArgumentParser()
parser.add_argument("--model_type", type=str, default='mdx23c',
help="One of mdx23c, htdemucs, segm_models, mel_band_roformer,"
" bs_roformer, swin_upernet, bandit")
parser.add_argument("--config_path", type=str, help="Path to config file")
parser.add_argument("--start_check_point", type=str, default='', help="Initial checkpoint"
" to valid weights")
parser.add_argument("--valid_path", nargs="+", type=str, help="Validate path")
parser.add_argument("--store_dir", type=str, default="", help="Path to store results as wav file")
parser.add_argument("--draw_spectro", type=float, default=0,
help="If --store_dir is set then code will generate spectrograms for resulted stems as well."
" Value defines for how many seconds os track spectrogram will be generated.")
parser.add_argument("--device_ids", nargs='+', type=int, default=[0], help='List of gpu ids')
parser.add_argument("--num_workers", type=int, default=0, help="Dataloader num_workers")
parser.add_argument("--pin_memory", action='store_true', help="Dataloader pin_memory")
parser.add_argument("--extension", type=str, default='wav', help="Choose extension for validation")
parser.add_argument("--use_tta", action='store_true',
help="Flag adds test time augmentation during inference (polarity and channel inverse)."
"While this triples the runtime, it reduces noise and slightly improves prediction quality.")
parser.add_argument("--metrics", nargs='+', type=str, default=["sdr"],
choices=['sdr', 'l1_freq', 'si_sdr', 'neg_log_wmse', 'aura_stft', 'aura_mrstft', 'bleedless',
'fullness'], help='List of metrics to use.')
parser.add_argument("--lora_checkpoint", type=str, default='', help="Initial checkpoint to LoRA weights")
if dict_args is not None:
args = parser.parse_args([])
args_dict = vars(args)
args_dict.update(dict_args)
args = argparse.Namespace(**args_dict)
else:
args = parser.parse_args()
return args
def parse_args_inference(dict_args: Union[Dict, None]) -> argparse.Namespace:
"""
Parse command-line arguments for inference configuration.
Builds the CLI for model selection, configuration path, input/output handling,
device/runtime options, test-time augmentation, and optional LoRA checkpoints.
If `dict_args` is provided, its key–value pairs override or supply CLI options
programmatically; otherwise, arguments are read from `sys.argv`.
Args:
dict_args (Union[Dict, None]): Optional mapping of argument names to values
used to override or supply CLI options programmatically.
Returns:
argparse.Namespace: Parsed arguments namespace containing all inference
configuration values.
"""
parser = argparse.ArgumentParser()
parser.add_argument("--model_type", type=str, default='mdx23c',
help="One of bandit, bandit_v2, bs_roformer, htdemucs, mdx23c, mel_band_roformer,"
" scnet, scnet_unofficial, segm_models, swin_upernet, torchseg")
parser.add_argument("--config_path", type=str, help="path to config file")
parser.add_argument("--start_check_point", type=str, default='', help="Initial checkpoint to valid weights")
parser.add_argument("--input_folder", type=str, help="folder with mixtures to process")
parser.add_argument("--store_dir", type=str, default="", help="path to store results as wav file")
parser.add_argument("--draw_spectro", type=float, default=0,
help="Code will generate spectrograms for resulted stems."
" Value defines for how many seconds os track spectrogram will be generated.")
parser.add_argument("--device_ids", nargs='+', type=int, default=0, help='list of gpu ids')
parser.add_argument("--extract_instrumental", action='store_true',
help="invert vocals to get instrumental if provided")
parser.add_argument("--disable_detailed_pbar", action='store_true', help="disable detailed progress bar")
parser.add_argument("--force_cpu", action='store_true', help="Force the use of CPU even if CUDA is available")
parser.add_argument("--flac_file", action='store_true', help="Output flac file instead of wav")
parser.add_argument("--pcm_type", type=str, choices=['PCM_16', 'PCM_24'], default='PCM_24',
help="PCM type for FLAC files (PCM_16 or PCM_24)")
parser.add_argument("--use_tta", action='store_true',
help="Flag adds test time augmentation during inference (polarity and channel inverse)."
"While this triples the runtime, it reduces noise and slightly improves prediction quality.")
parser.add_argument("--lora_checkpoint", type=str, default='', help="Initial checkpoint to LoRA weights")
if dict_args is not None:
args = parser.parse_args([])
args_dict = vars(args)
args_dict.update(dict_args)
args = argparse.Namespace(**args_dict)
else:
args = parser.parse_args()
return args
def load_config(model_type: str, config_path: str) -> Union[ConfigDict, OmegaConf]:
"""
Load a model configuration from a file.
Based on `model_type`, returns either an OmegaConf (e.g., for 'htdemucs')
or a YAML-parsed ConfigDict for other models.
Args:
model_type (str): Model identifier that determines the loader behavior
(e.g., 'htdemucs', 'mdx23c', etc.).
config_path (str): Path to the configuration file (YAML/OmegaConf).
Returns:
Union[ConfigDict, OmegaConf]: Loaded configuration object.
Raises:
FileNotFoundError: If `config_path` does not point to an existing file.
ValueError: If the configuration cannot be parsed or is otherwise invalid.
"""
try:
with open(config_path, 'r') as f:
if model_type == 'htdemucs':
config = OmegaConf.load(config_path)
else:
config = ConfigDict(yaml.load(f, Loader=yaml.FullLoader))
return config
except FileNotFoundError:
raise FileNotFoundError(f"Configuration file not found at {config_path}")
except Exception as e:
raise ValueError(f"Error loading configuration: {e}")
def get_model_from_config(model_type: str, config_path: str) -> Tuple[nn.Module, Union[ConfigDict, OmegaConf]]:
"""
Load and instantiate a model using a configuration file.
Given a `model_type` and a path to a configuration, this function loads the
configuration (YAML or OmegaConf) and constructs the corresponding model.
Args:
model_type (str): Identifier of the model family (e.g., 'mdx23c', 'htdemucs',
'scnet', 'mel_band_conformer', etc.).
config_path (str): Filesystem path to the configuration file used to
initialize the model.
Returns:
Tuple[nn.Module, Union[ConfigDict, OmegaConf]]: A tuple containing the
initialized PyTorch model and the loaded configuration object.
Raises:
ValueError: If `model_type` is unknown or model initialization fails.
FileNotFoundError: If `config_path` does not exist (may be raised by the
underlying config loader).
"""
config = load_config(model_type, config_path)
if model_type == 'mel_band_roformer':
from ..modules.bs_roformer import MelBandRoformer
model = MelBandRoformer(**dict(config.model))
else:
raise ValueError(f"Unknown model type: {model_type}")
return model, config
def logging(logs: List[str], text: str, verbose_logging: bool = False) -> None:
"""
Print a log message and optionally append it to an in-memory list.
In Distributed Data Parallel (DDP) contexts, the message is printed only on
rank 0; when DDP is uninitialized, it prints unconditionally. If
`verbose_logging` is True, the message is also appended to `logs`.
Args:
logs (List[str]): Mutable list to which the message is appended when
`verbose_logging` is True.
text (str): The log message to print (rank 0 only under DDP) and
optionally store.
verbose_logging (bool, optional): If True, append `text` to `logs`.
Defaults to False.
Returns:
None: The function prints and may mutate `logs` in place.
"""
if not dist.is_initialized() or dist.get_rank()==0:
print(text)
if verbose_logging:
logs.append(text)
def write_results_in_file(store_dir: str, logs: List[str]) -> None:
"""
Write accumulated log messages to a results file.
Creates (or overwrites) a `results.txt` file inside `store_dir` and writes
each entry from `logs` as a separate line. In Distributed Data Parallel (DDP)
scenarios, writing is intended to occur only on rank 0.
Args:
store_dir (str): Directory path where `results.txt` will be saved.
logs (List[str]): Ordered collection of log lines to write.
Returns:
None
"""
if not dist.is_initialized() or dist.get_rank() == 0:
with open(f'{store_dir}/results.txt', 'w') as out:
for item in logs:
out.write(item + "\n")
def manual_seed(seed: int) -> None:
"""
Initialize random seeds for reproducibility.
Sets the seed across Python's `random`, NumPy, and PyTorch (CPU and CUDA)
libraries, and updates the `PYTHONHASHSEED` environment variable. This helps
ensure deterministic behavior where possible, though some GPU operations
may still introduce nondeterminism.
Args:
seed (int): The seed value to use for all random number generators.
Returns:
None
"""
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
torch.cuda.manual_seed_all(seed) # if multi-GPU
torch.backends.cudnn.deterministic = False
os.environ["PYTHONHASHSEED"] = str(seed)
def initialize_environment(seed: int, results_path: str) -> None:
"""
Initialize runtime environment settings.
Sets random seeds for reproducibility, adjusts PyTorch cuDNN behavior,
configures multiprocessing with the 'spawn' start method, and ensures
the results directory exists.
Args:
seed (int): Random seed value for deterministic initialization.
results_path (str): Filesystem path to create for saving results.
Returns:
None
"""
manual_seed(seed)
torch.backends.cudnn.deterministic = False
try:
torch.multiprocessing.set_start_method('spawn')
except Exception as e:
pass
os.makedirs(results_path, exist_ok=True)
def initialize_environment_ddp(rank: int, world_size: int, seed: int = 0, resuls_path: str = None) -> None:
"""
Initialize environment for Distributed Data Parallel (DDP) training/validation.
Sets up the DDP process group, seeds random number generators, configures
multiprocessing to use the 'spawn' method, and creates a results directory
if provided.
Args:
rank (int): Rank of the current process within the DDP group.
world_size (int): Total number of processes participating in DDP.
seed (int, optional): Random seed for reproducibility. Defaults to 0.
resuls_path (str, optional): Directory path to create for storing results.
If None, no directory is created. Defaults to None.
Returns:
None
"""
setup_ddp(rank, world_size)
manual_seed(seed)
try:
torch.multiprocessing.set_start_method('spawn', force=True) # force=True prevent errors
except RuntimeError as e:
if "context has already been set" not in str(e):
raise e
if not(resuls_path is None):
os.makedirs(resuls_path, exist_ok=True)
def gen_wandb_name(args, config) -> str:
"""
Generate a descriptive name for a Weights & Biases (wandb) run.
Combines the model type, a dash-joined list of training instruments,
and the current date into a single string identifier.
Args:
args: Parsed arguments namespace containing at least `model_type`.
config: Configuration object/dict with a `training.instruments` field.
Returns:
str: Formatted run name in the form
"<model_type>_[<instrument1>-<instrument2>-...]_<YYYY-MM-DD>".
"""
instrum = '-'.join(config['training']['instruments'])
time_str = time.strftime("%Y-%m-%d")
name = '{}_[{}]_{}'.format(args.model_type, instrum, time_str)
return name
def wandb_init(args: argparse.Namespace, config: Dict, batch_size: int) -> None:
"""
Initialize Weights & Biases (wandb) for experiment tracking.
Depending on the provided arguments, sets up wandb in one of three modes:
- Offline mode when `args.wandb_offline` is True.
- Disabled mode when no valid `wandb_key` is provided.
- Online mode with authentication using `args.wandb_key`.
Args:
args (argparse.Namespace): Parsed arguments containing wandb options
(`wandb_offline`, `wandb_key`, `device_ids`).
config (Dict): Experiment configuration dictionary to log.
batch_size (int): Training batch size to include in the run configuration.
Returns:
None
"""
if not dist.is_initialized() or dist.get_rank() == 0:
if args.wandb_offline:
wandb.init(mode='offline',
project='msst',
name=gen_wandb_name(args, config),
config={'config': config, 'args': args, 'device_ids': args.device_ids, 'batch_size': batch_size}
)
elif args.wandb_key is None or args.wandb_key.strip() == '':
wandb.init(mode='disabled')
else:
wandb.login(key=args.wandb_key)
wandb.init(
project='msst',
name=gen_wandb_name(args, config),
config={'config': config, 'args': args, 'device_ids': args.device_ids, 'batch_size': batch_size}
)
def setup_ddp(rank: int, world_size: int) -> None:
"""
Initialize a Distributed Data Parallel (DDP) process group.
Configures environment variables for the DDP master node, attempts to
initialize the process group with the NCCL backend (preferred for GPUs),
and falls back to the Gloo backend if NCCL is unavailable. Also sets the
current CUDA device to match the process rank.
Args:
rank (int): Rank of the current process in the DDP group.
world_size (int): Total number of processes participating in DDP.
Returns:
None
"""
os.environ['MASTER_ADDR'] = 'localhost'
os.environ['MASTER_PORT'] = '12355' # We can change and use another
os.environ["USE_LIBUV"] = "0"
try:
dist.init_process_group("nccl", rank=rank, world_size=world_size)
except:
dist.init_process_group("gloo", rank=rank, world_size=world_size)
if dist.get_rank()==0:
print(f'NCCL are not available. Using "gloo" backend.')
torch.cuda.set_device(rank)
def cleanup_ddp() -> None:
"""
Finalize and clean up a Distributed Data Parallel (DDP) process group.
Calls `torch.distributed.destroy_process_group()` to release resources
associated with the current DDP environment.
Returns:
None
"""
dist.destroy_process_group()