Initial commit
This commit is contained in:
@@ -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()
|
||||
Reference in New Issue
Block a user