Initial commit
This commit is contained in:
@@ -0,0 +1,225 @@
|
||||
# https://github.com/ZFTurbo/Music-Source-Separation-Training
|
||||
# https://huggingface.co/becruily/mel-band-roformer-karaoke/blob/main/mel_band_roformer_karaoke_becruily.ckpt
|
||||
# https://huggingface.co/anvuew/dereverb_mel_band_roformer/blob/main/dereverb_mel_band_roformer_anvuew_sdr_19.1729.ckpt
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, Optional, Tuple
|
||||
|
||||
import librosa
|
||||
import sys
|
||||
import os
|
||||
import time
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
from .utils.audio_utils import normalize_audio, denormalize_audio
|
||||
from .utils.settings import get_model_from_config, parse_args_inference
|
||||
from .utils.model_utils import demix
|
||||
from .utils.model_utils import prefer_target_instrument, apply_tta, load_start_checkpoint
|
||||
|
||||
|
||||
def process(mix, model, args, config, device):
|
||||
|
||||
instruments = prefer_target_instrument(config)[:]
|
||||
|
||||
# If mono audio we must adjust it depending on model
|
||||
if len(mix.shape) == 1:
|
||||
mix = np.expand_dims(mix, axis=0)
|
||||
if 'num_channels' in config.audio:
|
||||
if config.audio['num_channels'] == 2:
|
||||
# print(f'Convert mono track to stereo...')
|
||||
mix = np.concatenate([mix, mix], axis=0)
|
||||
|
||||
if 'normalize' in config.inference:
|
||||
if config.inference['normalize'] is True:
|
||||
mix, norm_params = normalize_audio(mix)
|
||||
|
||||
waveforms_orig = demix(config, model, mix, device, model_type=args.model_type, pbar=not args.disable_detailed_pbar)
|
||||
|
||||
instr = 'vocals' if 'vocals' in instruments else instruments[0]
|
||||
estimates = waveforms_orig[instr]
|
||||
if 'normalize' in config.inference:
|
||||
if config.inference['normalize'] is True:
|
||||
estimates = denormalize_audio(estimates, norm_params)
|
||||
|
||||
return estimates
|
||||
|
||||
|
||||
def build_model(args):
|
||||
model, config = get_model_from_config(args.model_type, args.config_path)
|
||||
|
||||
load_start_checkpoint(args, model, None, type_='inference')
|
||||
|
||||
return model, config
|
||||
|
||||
|
||||
def build_models(dict_args):
|
||||
args = parse_args_inference(dict_args)
|
||||
|
||||
########## load model ##########
|
||||
torch.backends.cudnn.benchmark = True
|
||||
|
||||
args.config_path = args.sep_config_path
|
||||
args.start_check_point = args.sep_start_check_point
|
||||
|
||||
sep_model, sep_config = build_model(args)
|
||||
|
||||
args.config_path = args.der_config_path
|
||||
args.start_check_point = args.der_start_check_point
|
||||
|
||||
dereverb_model, dereverb_config = build_model(args)
|
||||
|
||||
sep_model = sep_model
|
||||
dereverb_model = dereverb_model
|
||||
|
||||
return sep_model, sep_config, dereverb_model, dereverb_config, args
|
||||
|
||||
def main(args, sep_model=None, sep_config=None, dereverb_model=None, dereverb_config=None, device=None):
|
||||
|
||||
######## process data ##########
|
||||
sample_rate = getattr(sep_config.audio, 'sample_rate', 44100)
|
||||
path = args.input_path
|
||||
|
||||
mix, _ = librosa.load(path, sr=sample_rate, mono=False)
|
||||
vocals = process(mix, sep_model, args, sep_config, device)
|
||||
dereverbed_vocals = process(vocals.mean(0), dereverb_model, args, dereverb_config, device)
|
||||
accompaniment = mix - dereverbed_vocals
|
||||
|
||||
return mix, vocals, dereverbed_vocals, accompaniment, sample_rate
|
||||
|
||||
@dataclass
|
||||
class VocalSeparationOutputs:
|
||||
"""Vocal extraction output container."""
|
||||
|
||||
mix: np.ndarray
|
||||
vocals: np.ndarray
|
||||
vocals_dereverbed: np.ndarray
|
||||
accompaniment: np.ndarray
|
||||
sample_rate: int
|
||||
|
||||
|
||||
class VocalSeparator:
|
||||
"""Vocal separation and dereverb wrapper.
|
||||
|
||||
Wraps the karaoke separation and dereverb models from the
|
||||
ZFTurbo Music Source Separation project and exposes a simple
|
||||
:py:meth:`process` API that returns mix/vocals/dereverbed/accompaniment.
|
||||
"""
|
||||
def __init__(
|
||||
self,
|
||||
sep_model_path: str,
|
||||
sep_config_path: str,
|
||||
der_model_path: str,
|
||||
der_config_path: str,
|
||||
*,
|
||||
model_type: str = "mel_band_roformer",
|
||||
disable_detailed_pbar: bool = True,
|
||||
device: str = "cuda",
|
||||
verbose: bool = True,
|
||||
):
|
||||
"""Initialize the vocal separator.
|
||||
|
||||
Args:
|
||||
device: Torch device string, e.g. ``"cuda:0"``.
|
||||
model_type: Separation model type key.
|
||||
sep_config_path: Config path for separation model.
|
||||
sep_start_check_point: Checkpoint path for separation model.
|
||||
der_config_path: Config path for dereverb model.
|
||||
der_start_check_point: Checkpoint path for dereverb model.
|
||||
disable_detailed_pbar: Disable detailed progress bars in underlying utils.
|
||||
verbose: Whether to print verbose logs.
|
||||
"""
|
||||
|
||||
# Match original script args schema
|
||||
args_dict: Dict[str, Any] = {
|
||||
"model_type": model_type,
|
||||
"disable_detailed_pbar": disable_detailed_pbar,
|
||||
"sep_config_path": sep_config_path,
|
||||
"sep_start_check_point": sep_model_path,
|
||||
"der_config_path": der_config_path,
|
||||
"der_start_check_point": der_model_path,
|
||||
}
|
||||
|
||||
if verbose:
|
||||
print("[vocal extraction] init: start")
|
||||
|
||||
sep_model, sep_config, dereverb_model, dereverb_config, args = build_models(args_dict)
|
||||
|
||||
sep_model = sep_model.to(device)
|
||||
dereverb_model = dereverb_model.to(device)
|
||||
|
||||
self.sep_model = sep_model
|
||||
self.sep_config = sep_config
|
||||
self.dereverb_model = dereverb_model
|
||||
self.dereverb_config = dereverb_config
|
||||
self.device = device
|
||||
self.args = args
|
||||
self.verbose = verbose
|
||||
|
||||
if verbose:
|
||||
print(
|
||||
"[vocal extraction] init success: sep=loaded, dereverb=loaded, device=",
|
||||
device,
|
||||
)
|
||||
|
||||
def process(self, input_path: str, *, verbose: Optional[bool] = None) -> VocalSeparationOutputs:
|
||||
"""Separate a single audio file into sources.
|
||||
|
||||
Args:
|
||||
input_path: Path to the mixture wav.
|
||||
verbose: Override instance-level verbose flag for this call.
|
||||
|
||||
Returns:
|
||||
:class:`VocalSeparationOutputs` containing mix, vocals,
|
||||
dereverbed vocals, accompaniment and sample rate.
|
||||
"""
|
||||
verbose = self.verbose if verbose is None else verbose
|
||||
if verbose:
|
||||
print(f"[vocal extraction] process_file: start: {input_path}")
|
||||
t0 = time.time()
|
||||
|
||||
self.args.input_path = input_path
|
||||
|
||||
mix, vocals, dereverbed, accompaniment, sample_rate = main(
|
||||
self.args,
|
||||
self.sep_model,
|
||||
self.sep_config,
|
||||
self.dereverb_model,
|
||||
self.dereverb_config,
|
||||
torch.device(self.device) if not isinstance(self.device, torch.device) else self.device,
|
||||
)
|
||||
|
||||
if verbose:
|
||||
dt = time.time() - t0
|
||||
print(
|
||||
"[vocal extraction] process_file: done:",
|
||||
f"sr={sample_rate}",
|
||||
f"mix={getattr(mix, 'shape', None)}",
|
||||
f"vocals={getattr(vocals, 'shape', None)}",
|
||||
f"dereverbed={getattr(dereverbed, 'shape', None)}",
|
||||
f"acc={getattr(accompaniment, 'shape', None)}",
|
||||
f"time={dt:.3f}s",
|
||||
)
|
||||
|
||||
return VocalSeparationOutputs(
|
||||
mix=mix,
|
||||
vocals=vocals,
|
||||
vocals_dereverbed=dereverbed,
|
||||
accompaniment=accompaniment,
|
||||
sample_rate=sample_rate,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
m = VocalSeparator(
|
||||
sep_model_path="pretrained_models/mel-band-roformer-karaoke/mel_band_roformer_karaoke_becruily.ckpt",
|
||||
sep_config_path="pretrained_models/mel-band-roformer-karaoke/config_karaoke_becruily.yaml",
|
||||
der_model_path="pretrained_models/dereverb_mel_band_roformer/dereverb_mel_band_roformer_anvuew_sdr_19.1729.ckpt",
|
||||
der_config_path="pretrained_models/dereverb_mel_band_roformer/dereverb_mel_band_roformer_anvuew.yaml",
|
||||
device="cuda"
|
||||
)
|
||||
|
||||
out = m.process("example/test/separation_test.mp3")
|
||||
print(out.vocals_dereverbed.shape)
|
||||
@@ -0,0 +1,2 @@
|
||||
from .bs_roformer import BSRoformer
|
||||
from .mel_band_roformer import MelBandRoformer
|
||||
@@ -0,0 +1,126 @@
|
||||
from functools import wraps
|
||||
from packaging import version
|
||||
from collections import namedtuple
|
||||
|
||||
import os
|
||||
import torch
|
||||
from torch import nn, einsum
|
||||
import torch.nn.functional as F
|
||||
|
||||
from einops import rearrange, reduce
|
||||
|
||||
# constants
|
||||
|
||||
FlashAttentionConfig = namedtuple('FlashAttentionConfig', ['enable_flash', 'enable_math', 'enable_mem_efficient'])
|
||||
|
||||
# helpers
|
||||
|
||||
def exists(val):
|
||||
return val is not None
|
||||
|
||||
def default(v, d):
|
||||
return v if exists(v) else d
|
||||
|
||||
def once(fn):
|
||||
called = False
|
||||
@wraps(fn)
|
||||
def inner(x):
|
||||
nonlocal called
|
||||
if called:
|
||||
return
|
||||
called = True
|
||||
return fn(x)
|
||||
return inner
|
||||
|
||||
print_once = once(print)
|
||||
|
||||
# main class
|
||||
|
||||
class Attend(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dropout = 0.,
|
||||
flash = False,
|
||||
scale = None
|
||||
):
|
||||
super().__init__()
|
||||
self.scale = scale
|
||||
self.dropout = dropout
|
||||
self.attn_dropout = nn.Dropout(dropout)
|
||||
|
||||
self.flash = flash
|
||||
assert not (flash and version.parse(torch.__version__) < version.parse('2.0.0')), 'in order to use flash attention, you must be using pytorch 2.0 or above'
|
||||
|
||||
# determine efficient attention configs for cuda and cpu
|
||||
|
||||
self.cpu_config = FlashAttentionConfig(True, True, True)
|
||||
self.cuda_config = None
|
||||
|
||||
if not torch.cuda.is_available() or not flash:
|
||||
return
|
||||
|
||||
device_properties = torch.cuda.get_device_properties(torch.device('cuda'))
|
||||
device_version = version.parse(f'{device_properties.major}.{device_properties.minor}')
|
||||
|
||||
if device_version >= version.parse('8.0'):
|
||||
if os.name == 'nt':
|
||||
print_once('Windows OS detected, using math or mem efficient attention if input tensor is on cuda')
|
||||
self.cuda_config = FlashAttentionConfig(False, True, True)
|
||||
else:
|
||||
print_once('GPU Compute Capability equal or above 8.0, using flash attention if input tensor is on cuda')
|
||||
self.cuda_config = FlashAttentionConfig(True, False, False)
|
||||
else:
|
||||
print_once('GPU Compute Capability below 8.0, using math or mem efficient attention if input tensor is on cuda')
|
||||
self.cuda_config = FlashAttentionConfig(False, True, True)
|
||||
|
||||
def flash_attn(self, q, k, v):
|
||||
_, heads, q_len, _, k_len, is_cuda, device = *q.shape, k.shape[-2], q.is_cuda, q.device
|
||||
|
||||
if exists(self.scale):
|
||||
default_scale = q.shape[-1] ** -0.5
|
||||
q = q * (self.scale / default_scale)
|
||||
|
||||
# Check if there is a compatible device for flash attention
|
||||
|
||||
config = self.cuda_config if is_cuda else self.cpu_config
|
||||
|
||||
# pytorch 2.0 flash attn: q, k, v, mask, dropout, softmax_scale
|
||||
|
||||
with torch.backends.cuda.sdp_kernel(**config._asdict()):
|
||||
out = F.scaled_dot_product_attention(
|
||||
q, k, v,
|
||||
dropout_p = self.dropout if self.training else 0.
|
||||
)
|
||||
|
||||
return out
|
||||
|
||||
def forward(self, q, k, v):
|
||||
"""
|
||||
einstein notation
|
||||
b - batch
|
||||
h - heads
|
||||
n, i, j - sequence length (base sequence length, source, target)
|
||||
d - feature dimension
|
||||
"""
|
||||
|
||||
q_len, k_len, device = q.shape[-2], k.shape[-2], q.device
|
||||
|
||||
scale = default(self.scale, q.shape[-1] ** -0.5)
|
||||
|
||||
if self.flash:
|
||||
return self.flash_attn(q, k, v)
|
||||
|
||||
# similarity
|
||||
|
||||
sim = einsum(f"b h i d, b h j d -> b h i j", q, k) * scale
|
||||
|
||||
# attention
|
||||
|
||||
attn = sim.softmax(dim=-1)
|
||||
attn = self.attn_dropout(attn)
|
||||
|
||||
# aggregate values
|
||||
|
||||
out = einsum(f"b h i j, b h j d -> b h i d", attn, v)
|
||||
|
||||
return out
|
||||
@@ -0,0 +1,145 @@
|
||||
from functools import wraps
|
||||
from packaging import version
|
||||
from collections import namedtuple
|
||||
|
||||
import os
|
||||
import torch
|
||||
from torch import nn, einsum
|
||||
import torch.nn.functional as F
|
||||
|
||||
from einops import rearrange, reduce
|
||||
|
||||
def _print_once(msg):
|
||||
printed = False
|
||||
@wraps(print)
|
||||
def inner():
|
||||
nonlocal printed
|
||||
if not printed:
|
||||
print(msg)
|
||||
printed = True
|
||||
return inner
|
||||
|
||||
try:
|
||||
from sageattention import sageattn
|
||||
_has_sage_attention = True
|
||||
# _print_sage_found = _print_once("SageAttention found. Will be used when flash=True.")
|
||||
# _print_sage_found()
|
||||
except ImportError:
|
||||
_has_sage_attention = False
|
||||
_print_sage_not_found = _print_once("SageAttention not found. Will fall back to PyTorch SDPA (if available) or manual einsum.")
|
||||
_print_sage_not_found()
|
||||
|
||||
# helpers
|
||||
def exists(val):
|
||||
return val is not None
|
||||
|
||||
def default(v, d):
|
||||
return v if exists(v) else d
|
||||
|
||||
# main class
|
||||
class Attend(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dropout = 0.,
|
||||
flash = False, # If True, attempts to use SageAttention or PyTorch SDPA
|
||||
scale = None
|
||||
):
|
||||
super().__init__()
|
||||
self.scale = scale # Store the scale if needed for einsum path
|
||||
self.dropout = dropout # Store dropout if needed for einsum/SDPA path
|
||||
|
||||
# Determine which attention mechanism to *try* first
|
||||
self.use_sage = flash and _has_sage_attention
|
||||
self.use_pytorch_sdpa = False
|
||||
self._sdpa_checked = False # Flag to check PyTorch version only once
|
||||
|
||||
if flash and not self.use_sage:
|
||||
# Only consider PyTorch SDPA if Sage isn't available/chosen
|
||||
if not self._sdpa_checked:
|
||||
if version.parse(torch.__version__) >= version.parse('2.0.0'):
|
||||
self.use_pytorch_sdpa = True
|
||||
_print_sdpa_used = _print_once("Using PyTorch SDPA backend (FlashAttention-2, Memory-Efficient, or Math).")
|
||||
_print_sdpa_used()
|
||||
else:
|
||||
_print_fallback_einsum = _print_once("Flash attention requested but Pytorch < 2.0 and SageAttention not found. Falling back to einsum.")
|
||||
_print_fallback_einsum()
|
||||
self._sdpa_checked = True
|
||||
|
||||
# Dropout layer for manual einsum implementation ONLY
|
||||
# SDPA and SageAttention handle dropout differently (or not at all in Sage's base API)
|
||||
self.attn_dropout = nn.Dropout(dropout)
|
||||
|
||||
def forward(self, q, k, v):
|
||||
"""
|
||||
einstein notation
|
||||
b - batch
|
||||
h - heads
|
||||
n, i, j - sequence length (base sequence length, source, target)
|
||||
d - feature dimension
|
||||
|
||||
Input tensors q, k, v expected in shape: (batch, heads, seq_len, dim_head) -> HND layout
|
||||
"""
|
||||
q_len, k_len, device = q.shape[-2], k.shape[-2], q.device
|
||||
|
||||
# --- Priority 1: SageAttention ---
|
||||
if self.use_sage:
|
||||
# Assumes q, k, v are FP16/BF16 (handled by autocast upstream)
|
||||
# Assumes scale is handled internally by sageattn
|
||||
# Assumes dropout is NOT handled by sageattn kernel
|
||||
# is_causal=False based on how Attend is called in mel_band_roformer
|
||||
out = sageattn(q, k, v, tensor_layout='HND', is_causal=False)
|
||||
return out
|
||||
try:
|
||||
return out
|
||||
# print("Attempting SageAttention") # Optional: for debugging
|
||||
out = sageattn(q, k, v, tensor_layout='HND', is_causal=False)
|
||||
return out
|
||||
except Exception as e:
|
||||
print(f"SageAttention failed with error: {e}. Falling back.")
|
||||
self.use_sage = False # Don't try Sage again if it failed once
|
||||
# Decide fallback: Check if PyTorch SDPA is an option
|
||||
if not self._sdpa_checked:
|
||||
if version.parse(torch.__version__) >= version.parse('2.0.0'):
|
||||
self.use_pytorch_sdpa = True
|
||||
_print_sdpa_fallback = _print_once("Falling back to PyTorch SDPA.")
|
||||
_print_sdpa_fallback()
|
||||
else:
|
||||
_print_einsum_fallback = _print_once("Falling back to einsum.")
|
||||
_print_einsum_fallback()
|
||||
self._sdpa_checked = True
|
||||
|
||||
|
||||
# --- Priority 2: PyTorch SDPA ---
|
||||
if self.use_pytorch_sdpa:
|
||||
# Use PyTorch's Scaled Dot Product Attention (SDPA)
|
||||
# It handles scaling and dropout internally.
|
||||
try:
|
||||
# print("Attempting PyTorch SDPA") # Optional: for debugging
|
||||
# Let PyTorch choose the best backend (Flash V2, Mem Efficient, Math)
|
||||
with torch.backends.cuda.sdp_kernel(enable_flash=True, enable_math=True, enable_mem_efficient=True):
|
||||
out = F.scaled_dot_product_attention(
|
||||
q, k, v,
|
||||
attn_mask=None, # Assuming no explicit mask needed here
|
||||
dropout_p = self.dropout if self.training else 0.,
|
||||
is_causal=False # Assuming not needed based on usage context
|
||||
)
|
||||
return out
|
||||
except Exception as e:
|
||||
print(f"PyTorch SDPA failed with error: {e}. Falling back to einsum.")
|
||||
self.use_pytorch_sdpa = False # Fallback to einsum on error
|
||||
|
||||
|
||||
# Calculate scale
|
||||
scale = default(self.scale, q.shape[-1] ** -0.5)
|
||||
|
||||
# similarity
|
||||
sim = einsum(f"b h i d, b h j d -> b h i j", q, k) * scale
|
||||
|
||||
# attention
|
||||
attn = sim.softmax(dim=-1)
|
||||
attn = self.attn_dropout(attn) # Apply dropout ONLY in einsum path
|
||||
|
||||
# aggregate values
|
||||
out = einsum(f"b h i j, b h j d -> b h i d", attn, v)
|
||||
|
||||
return out
|
||||
@@ -0,0 +1,658 @@
|
||||
from functools import partial
|
||||
|
||||
import torch
|
||||
from torch import nn, einsum, Tensor
|
||||
from torch.nn import Module, ModuleList
|
||||
import torch.nn.functional as F
|
||||
|
||||
from .attend import Attend
|
||||
try:
|
||||
from .attend_sage import Attend as AttendSage
|
||||
except:
|
||||
pass
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
from beartype.typing import Tuple, Optional, List, Callable
|
||||
from beartype import beartype
|
||||
|
||||
from rotary_embedding_torch import RotaryEmbedding
|
||||
|
||||
from einops import rearrange, pack, unpack
|
||||
from einops.layers.torch import Rearrange
|
||||
|
||||
# helper functions
|
||||
|
||||
def exists(val):
|
||||
return val is not None
|
||||
|
||||
|
||||
def default(v, d):
|
||||
return v if exists(v) else d
|
||||
|
||||
|
||||
def pack_one(t, pattern):
|
||||
return pack([t], pattern)
|
||||
|
||||
|
||||
def unpack_one(t, ps, pattern):
|
||||
return unpack(t, ps, pattern)[0]
|
||||
|
||||
|
||||
# norm
|
||||
|
||||
def l2norm(t):
|
||||
return F.normalize(t, dim = -1, p = 2)
|
||||
|
||||
|
||||
class RMSNorm(Module):
|
||||
def __init__(self, dim):
|
||||
super().__init__()
|
||||
self.scale = dim ** 0.5
|
||||
self.gamma = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def forward(self, x):
|
||||
return F.normalize(x, dim=-1) * self.scale * self.gamma
|
||||
|
||||
|
||||
# attention
|
||||
|
||||
class FeedForward(Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
mult=4,
|
||||
dropout=0.
|
||||
):
|
||||
super().__init__()
|
||||
dim_inner = int(dim * mult)
|
||||
self.net = nn.Sequential(
|
||||
RMSNorm(dim),
|
||||
nn.Linear(dim, dim_inner),
|
||||
nn.GELU(),
|
||||
nn.Dropout(dropout),
|
||||
nn.Linear(dim_inner, dim),
|
||||
nn.Dropout(dropout)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
return self.net(x)
|
||||
|
||||
|
||||
class Attention(Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
heads=8,
|
||||
dim_head=64,
|
||||
dropout=0.,
|
||||
rotary_embed=None,
|
||||
flash=True,
|
||||
sage_attention=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.heads = heads
|
||||
self.scale = dim_head ** -0.5
|
||||
dim_inner = heads * dim_head
|
||||
|
||||
self.rotary_embed = rotary_embed
|
||||
|
||||
if sage_attention:
|
||||
self.attend = AttendSage(flash=flash, dropout=dropout)
|
||||
else:
|
||||
self.attend = Attend(flash=flash, dropout=dropout)
|
||||
|
||||
self.norm = RMSNorm(dim)
|
||||
self.to_qkv = nn.Linear(dim, dim_inner * 3, bias=False)
|
||||
|
||||
self.to_gates = nn.Linear(dim, heads)
|
||||
|
||||
self.to_out = nn.Sequential(
|
||||
nn.Linear(dim_inner, dim, bias=False),
|
||||
nn.Dropout(dropout)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.norm(x)
|
||||
|
||||
q, k, v = rearrange(self.to_qkv(x), 'b n (qkv h d) -> qkv b h n d', qkv=3, h=self.heads)
|
||||
|
||||
if exists(self.rotary_embed):
|
||||
q = self.rotary_embed.rotate_queries_or_keys(q)
|
||||
k = self.rotary_embed.rotate_queries_or_keys(k)
|
||||
|
||||
out = self.attend(q, k, v)
|
||||
|
||||
gates = self.to_gates(x)
|
||||
out = out * rearrange(gates, 'b n h -> b h n 1').sigmoid()
|
||||
|
||||
out = rearrange(out, 'b h n d -> b n (h d)')
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class LinearAttention(Module):
|
||||
"""
|
||||
this flavor of linear attention proposed in https://arxiv.org/abs/2106.09681 by El-Nouby et al.
|
||||
"""
|
||||
|
||||
@beartype
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
dim,
|
||||
dim_head=32,
|
||||
heads=8,
|
||||
scale=8,
|
||||
flash=False,
|
||||
dropout=0.,
|
||||
sage_attention=False,
|
||||
):
|
||||
super().__init__()
|
||||
dim_inner = dim_head * heads
|
||||
self.norm = RMSNorm(dim)
|
||||
|
||||
self.to_qkv = nn.Sequential(
|
||||
nn.Linear(dim, dim_inner * 3, bias=False),
|
||||
Rearrange('b n (qkv h d) -> qkv b h d n', qkv=3, h=heads)
|
||||
)
|
||||
|
||||
self.temperature = nn.Parameter(torch.ones(heads, 1, 1))
|
||||
|
||||
if sage_attention:
|
||||
self.attend = AttendSage(
|
||||
scale=scale,
|
||||
dropout=dropout,
|
||||
flash=flash
|
||||
)
|
||||
else:
|
||||
self.attend = Attend(
|
||||
scale=scale,
|
||||
dropout=dropout,
|
||||
flash=flash
|
||||
)
|
||||
|
||||
self.to_out = nn.Sequential(
|
||||
Rearrange('b h d n -> b n (h d)'),
|
||||
nn.Linear(dim_inner, dim, bias=False)
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x
|
||||
):
|
||||
x = self.norm(x)
|
||||
|
||||
q, k, v = self.to_qkv(x)
|
||||
|
||||
q, k = map(l2norm, (q, k))
|
||||
q = q * self.temperature.exp()
|
||||
|
||||
out = self.attend(q, k, v)
|
||||
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class Transformer(Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
dim,
|
||||
depth,
|
||||
dim_head=64,
|
||||
heads=8,
|
||||
attn_dropout=0.,
|
||||
ff_dropout=0.,
|
||||
ff_mult=4,
|
||||
norm_output=True,
|
||||
rotary_embed=None,
|
||||
flash_attn=True,
|
||||
linear_attn=False,
|
||||
sage_attention=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.layers = ModuleList([])
|
||||
|
||||
for _ in range(depth):
|
||||
if linear_attn:
|
||||
attn = LinearAttention(
|
||||
dim=dim,
|
||||
dim_head=dim_head,
|
||||
heads=heads,
|
||||
dropout=attn_dropout,
|
||||
flash=flash_attn,
|
||||
sage_attention=sage_attention
|
||||
)
|
||||
else:
|
||||
attn = Attention(
|
||||
dim=dim,
|
||||
dim_head=dim_head,
|
||||
heads=heads,
|
||||
dropout=attn_dropout,
|
||||
rotary_embed=rotary_embed,
|
||||
flash=flash_attn,
|
||||
sage_attention=sage_attention
|
||||
)
|
||||
|
||||
self.layers.append(ModuleList([
|
||||
attn,
|
||||
FeedForward(dim=dim, mult=ff_mult, dropout=ff_dropout)
|
||||
]))
|
||||
|
||||
self.norm = RMSNorm(dim) if norm_output else nn.Identity()
|
||||
|
||||
def forward(self, x):
|
||||
|
||||
for attn, ff in self.layers:
|
||||
x = attn(x) + x
|
||||
x = ff(x) + x
|
||||
|
||||
return self.norm(x)
|
||||
|
||||
|
||||
# bandsplit module
|
||||
|
||||
class BandSplit(Module):
|
||||
@beartype
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
dim_inputs: Tuple[int, ...]
|
||||
):
|
||||
super().__init__()
|
||||
self.dim_inputs = dim_inputs
|
||||
self.to_features = ModuleList([])
|
||||
|
||||
for dim_in in dim_inputs:
|
||||
net = nn.Sequential(
|
||||
RMSNorm(dim_in),
|
||||
nn.Linear(dim_in, dim)
|
||||
)
|
||||
|
||||
self.to_features.append(net)
|
||||
|
||||
def forward(self, x):
|
||||
x = x.split(self.dim_inputs, dim=-1)
|
||||
|
||||
outs = []
|
||||
for split_input, to_feature in zip(x, self.to_features):
|
||||
split_output = to_feature(split_input)
|
||||
outs.append(split_output)
|
||||
|
||||
return torch.stack(outs, dim=-2)
|
||||
|
||||
|
||||
def MLP(
|
||||
dim_in,
|
||||
dim_out,
|
||||
dim_hidden=None,
|
||||
depth=1,
|
||||
activation=nn.Tanh
|
||||
):
|
||||
dim_hidden = default(dim_hidden, dim_in)
|
||||
|
||||
net = []
|
||||
dims = (dim_in, *((dim_hidden,) * (depth - 1)), dim_out)
|
||||
|
||||
for ind, (layer_dim_in, layer_dim_out) in enumerate(zip(dims[:-1], dims[1:])):
|
||||
is_last = ind == (len(dims) - 2)
|
||||
|
||||
net.append(nn.Linear(layer_dim_in, layer_dim_out))
|
||||
|
||||
if is_last:
|
||||
continue
|
||||
|
||||
net.append(activation())
|
||||
|
||||
return nn.Sequential(*net)
|
||||
|
||||
|
||||
class MaskEstimator(Module):
|
||||
@beartype
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
dim_inputs: Tuple[int, ...],
|
||||
depth,
|
||||
mlp_expansion_factor=4
|
||||
):
|
||||
super().__init__()
|
||||
self.dim_inputs = dim_inputs
|
||||
self.to_freqs = ModuleList([])
|
||||
dim_hidden = dim * mlp_expansion_factor
|
||||
|
||||
for dim_in in dim_inputs:
|
||||
net = []
|
||||
|
||||
mlp = nn.Sequential(
|
||||
MLP(dim, dim_in * 2, dim_hidden=dim_hidden, depth=depth),
|
||||
nn.GLU(dim=-1)
|
||||
)
|
||||
|
||||
self.to_freqs.append(mlp)
|
||||
|
||||
def forward(self, x):
|
||||
x = x.unbind(dim=-2)
|
||||
|
||||
outs = []
|
||||
|
||||
for band_features, mlp in zip(x, self.to_freqs):
|
||||
freq_out = mlp(band_features)
|
||||
outs.append(freq_out)
|
||||
|
||||
return torch.cat(outs, dim=-1)
|
||||
|
||||
|
||||
# main class
|
||||
|
||||
DEFAULT_FREQS_PER_BANDS = (
|
||||
2, 2, 2, 2, 2, 2, 2, 2, 2, 2,
|
||||
2, 2, 2, 2, 2, 2, 2, 2, 2, 2,
|
||||
2, 2, 2, 2,
|
||||
4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4,
|
||||
12, 12, 12, 12, 12, 12, 12, 12,
|
||||
24, 24, 24, 24, 24, 24, 24, 24,
|
||||
48, 48, 48, 48, 48, 48, 48, 48,
|
||||
128, 129,
|
||||
)
|
||||
|
||||
|
||||
class BSRoformer(Module):
|
||||
|
||||
@beartype
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
*,
|
||||
depth,
|
||||
stereo=False,
|
||||
num_stems=1,
|
||||
time_transformer_depth=2,
|
||||
freq_transformer_depth=2,
|
||||
linear_transformer_depth=0,
|
||||
freqs_per_bands: Tuple[int, ...] = DEFAULT_FREQS_PER_BANDS,
|
||||
# in the paper, they divide into ~60 bands, test with 1 for starters
|
||||
dim_head=64,
|
||||
heads=8,
|
||||
attn_dropout=0.,
|
||||
ff_dropout=0.,
|
||||
flash_attn=True,
|
||||
dim_freqs_in=1025,
|
||||
stft_n_fft=2048,
|
||||
stft_hop_length=512,
|
||||
# 10ms at 44100Hz, from sections 4.1, 4.4 in the paper - @faroit recommends // 2 or // 4 for better reconstruction
|
||||
stft_win_length=2048,
|
||||
stft_normalized=False,
|
||||
stft_window_fn: Optional[Callable] = None,
|
||||
mask_estimator_depth=2,
|
||||
multi_stft_resolution_loss_weight=1.,
|
||||
multi_stft_resolutions_window_sizes: Tuple[int, ...] = (4096, 2048, 1024, 512, 256),
|
||||
multi_stft_hop_size=147,
|
||||
multi_stft_normalized=False,
|
||||
multi_stft_window_fn: Callable = torch.hann_window,
|
||||
mlp_expansion_factor=4,
|
||||
use_torch_checkpoint=False,
|
||||
skip_connection=False,
|
||||
sage_attention=False,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.stereo = stereo
|
||||
self.audio_channels = 2 if stereo else 1
|
||||
self.num_stems = num_stems
|
||||
self.use_torch_checkpoint = use_torch_checkpoint
|
||||
self.skip_connection = skip_connection
|
||||
|
||||
self.layers = ModuleList([])
|
||||
|
||||
if sage_attention:
|
||||
print("Use Sage Attention")
|
||||
|
||||
transformer_kwargs = dict(
|
||||
dim=dim,
|
||||
heads=heads,
|
||||
dim_head=dim_head,
|
||||
attn_dropout=attn_dropout,
|
||||
ff_dropout=ff_dropout,
|
||||
flash_attn=flash_attn,
|
||||
norm_output=False,
|
||||
sage_attention=sage_attention,
|
||||
)
|
||||
|
||||
time_rotary_embed = RotaryEmbedding(dim=dim_head)
|
||||
freq_rotary_embed = RotaryEmbedding(dim=dim_head)
|
||||
|
||||
for _ in range(depth):
|
||||
tran_modules = []
|
||||
if linear_transformer_depth > 0:
|
||||
tran_modules.append(Transformer(depth=linear_transformer_depth, linear_attn=True, **transformer_kwargs))
|
||||
tran_modules.append(
|
||||
Transformer(depth=time_transformer_depth, rotary_embed=time_rotary_embed, **transformer_kwargs)
|
||||
)
|
||||
tran_modules.append(
|
||||
Transformer(depth=freq_transformer_depth, rotary_embed=freq_rotary_embed, **transformer_kwargs)
|
||||
)
|
||||
self.layers.append(nn.ModuleList(tran_modules))
|
||||
|
||||
self.final_norm = RMSNorm(dim)
|
||||
|
||||
self.stft_kwargs = dict(
|
||||
n_fft=stft_n_fft,
|
||||
hop_length=stft_hop_length,
|
||||
win_length=stft_win_length,
|
||||
normalized=stft_normalized
|
||||
)
|
||||
|
||||
self.stft_window_fn = partial(default(stft_window_fn, torch.hann_window), stft_win_length)
|
||||
|
||||
freqs = torch.stft(torch.randn(1, 4096), **self.stft_kwargs, window=torch.ones(stft_win_length), return_complex=True).shape[1]
|
||||
|
||||
assert len(freqs_per_bands) > 1
|
||||
assert sum(
|
||||
freqs_per_bands) == freqs, f'the number of freqs in the bands must equal {freqs} based on the STFT settings, but got {sum(freqs_per_bands)}'
|
||||
|
||||
freqs_per_bands_with_complex = tuple(2 * f * self.audio_channels for f in freqs_per_bands)
|
||||
|
||||
self.band_split = BandSplit(
|
||||
dim=dim,
|
||||
dim_inputs=freqs_per_bands_with_complex
|
||||
)
|
||||
|
||||
self.mask_estimators = nn.ModuleList([])
|
||||
|
||||
for _ in range(num_stems):
|
||||
mask_estimator = MaskEstimator(
|
||||
dim=dim,
|
||||
dim_inputs=freqs_per_bands_with_complex,
|
||||
depth=mask_estimator_depth,
|
||||
mlp_expansion_factor=mlp_expansion_factor,
|
||||
)
|
||||
|
||||
self.mask_estimators.append(mask_estimator)
|
||||
|
||||
# for the multi-resolution stft loss
|
||||
|
||||
self.multi_stft_resolution_loss_weight = multi_stft_resolution_loss_weight
|
||||
self.multi_stft_resolutions_window_sizes = multi_stft_resolutions_window_sizes
|
||||
self.multi_stft_n_fft = stft_n_fft
|
||||
self.multi_stft_window_fn = multi_stft_window_fn
|
||||
|
||||
self.multi_stft_kwargs = dict(
|
||||
hop_length=multi_stft_hop_size,
|
||||
normalized=multi_stft_normalized
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
raw_audio,
|
||||
target=None,
|
||||
return_loss_breakdown=False
|
||||
):
|
||||
"""
|
||||
einops
|
||||
|
||||
b - batch
|
||||
f - freq
|
||||
t - time
|
||||
s - audio channel (1 for mono, 2 for stereo)
|
||||
n - number of 'stems'
|
||||
c - complex (2)
|
||||
d - feature dimension
|
||||
"""
|
||||
|
||||
device = raw_audio.device
|
||||
|
||||
# defining whether model is loaded on MPS (MacOS GPU accelerator)
|
||||
x_is_mps = True if device.type == "mps" else False
|
||||
|
||||
if raw_audio.ndim == 2:
|
||||
raw_audio = rearrange(raw_audio, 'b t -> b 1 t')
|
||||
|
||||
channels = raw_audio.shape[1]
|
||||
assert (not self.stereo and channels == 1) or (self.stereo and channels == 2), 'stereo needs to be set to True if passing in audio signal that is stereo (channel dimension of 2). also need to be False if mono (channel dimension of 1)'
|
||||
|
||||
# to stft
|
||||
|
||||
raw_audio, batch_audio_channel_packed_shape = pack_one(raw_audio, '* t')
|
||||
|
||||
stft_window = self.stft_window_fn(device=device)
|
||||
|
||||
# RuntimeError: FFT operations are only supported on MacOS 14+
|
||||
# Since it's tedious to define whether we're on correct MacOS version - simple try-catch is used
|
||||
try:
|
||||
stft_repr = torch.stft(raw_audio, **self.stft_kwargs, window=stft_window, return_complex=True)
|
||||
except:
|
||||
stft_repr = torch.stft(raw_audio.cpu() if x_is_mps else raw_audio, **self.stft_kwargs,
|
||||
window=stft_window.cpu() if x_is_mps else stft_window, return_complex=True).to(
|
||||
device)
|
||||
stft_repr = torch.view_as_real(stft_repr)
|
||||
|
||||
stft_repr = unpack_one(stft_repr, batch_audio_channel_packed_shape, '* f t c')
|
||||
|
||||
# merge stereo / mono into the frequency, with frequency leading dimension, for band splitting
|
||||
stft_repr = rearrange(stft_repr,'b s f t c -> b (f s) t c')
|
||||
|
||||
x = rearrange(stft_repr, 'b f t c -> b t (f c)')
|
||||
|
||||
if self.use_torch_checkpoint:
|
||||
x = checkpoint(self.band_split, x, use_reentrant=False)
|
||||
else:
|
||||
x = self.band_split(x)
|
||||
|
||||
# axial / hierarchical attention
|
||||
|
||||
store = [None] * len(self.layers)
|
||||
for i, transformer_block in enumerate(self.layers):
|
||||
|
||||
if len(transformer_block) == 3:
|
||||
linear_transformer, time_transformer, freq_transformer = transformer_block
|
||||
|
||||
x, ft_ps = pack([x], 'b * d')
|
||||
if self.use_torch_checkpoint:
|
||||
x = checkpoint(linear_transformer, x, use_reentrant=False)
|
||||
else:
|
||||
x = linear_transformer(x)
|
||||
x, = unpack(x, ft_ps, 'b * d')
|
||||
else:
|
||||
time_transformer, freq_transformer = transformer_block
|
||||
|
||||
if self.skip_connection:
|
||||
# Sum all previous
|
||||
for j in range(i):
|
||||
x = x + store[j]
|
||||
|
||||
x = rearrange(x, 'b t f d -> b f t d')
|
||||
x, ps = pack([x], '* t d')
|
||||
|
||||
if self.use_torch_checkpoint:
|
||||
x = checkpoint(time_transformer, x, use_reentrant=False)
|
||||
else:
|
||||
x = time_transformer(x)
|
||||
|
||||
x, = unpack(x, ps, '* t d')
|
||||
x = rearrange(x, 'b f t d -> b t f d')
|
||||
x, ps = pack([x], '* f d')
|
||||
|
||||
if self.use_torch_checkpoint:
|
||||
x = checkpoint(freq_transformer, x, use_reentrant=False)
|
||||
else:
|
||||
x = freq_transformer(x)
|
||||
|
||||
x, = unpack(x, ps, '* f d')
|
||||
|
||||
if self.skip_connection:
|
||||
store[i] = x
|
||||
|
||||
x = self.final_norm(x)
|
||||
|
||||
num_stems = len(self.mask_estimators)
|
||||
|
||||
if self.use_torch_checkpoint:
|
||||
mask = torch.stack([checkpoint(fn, x, use_reentrant=False) for fn in self.mask_estimators], dim=1)
|
||||
else:
|
||||
mask = torch.stack([fn(x) for fn in self.mask_estimators], dim=1)
|
||||
mask = rearrange(mask, 'b n t (f c) -> b n f t c', c=2)
|
||||
|
||||
# modulate frequency representation
|
||||
|
||||
stft_repr = rearrange(stft_repr, 'b f t c -> b 1 f t c')
|
||||
|
||||
# complex number multiplication
|
||||
|
||||
stft_repr = torch.view_as_complex(stft_repr)
|
||||
mask = torch.view_as_complex(mask)
|
||||
|
||||
stft_repr = stft_repr * mask
|
||||
|
||||
# istft
|
||||
|
||||
stft_repr = rearrange(stft_repr, 'b n (f s) t -> (b n s) f t', s=self.audio_channels)
|
||||
|
||||
# same as torch.stft() fix for MacOS MPS above
|
||||
try:
|
||||
recon_audio = torch.istft(stft_repr, **self.stft_kwargs, window=stft_window, return_complex=False, length=raw_audio.shape[-1])
|
||||
except:
|
||||
recon_audio = torch.istft(stft_repr.cpu() if x_is_mps else stft_repr, **self.stft_kwargs, window=stft_window.cpu() if x_is_mps else stft_window, return_complex=False, length=raw_audio.shape[-1]).to(device)
|
||||
|
||||
recon_audio = rearrange(recon_audio, '(b n s) t -> b n s t', s=self.audio_channels, n=num_stems)
|
||||
|
||||
if num_stems == 1:
|
||||
recon_audio = rearrange(recon_audio, 'b 1 s t -> b s t')
|
||||
|
||||
# if a target is passed in, calculate loss for learning
|
||||
|
||||
if not exists(target):
|
||||
return recon_audio
|
||||
|
||||
if self.num_stems > 1:
|
||||
assert target.ndim == 4 and target.shape[1] == self.num_stems
|
||||
|
||||
if target.ndim == 2:
|
||||
target = rearrange(target, '... t -> ... 1 t')
|
||||
|
||||
target = target[..., :recon_audio.shape[-1]] # protect against lost length on istft
|
||||
|
||||
loss = F.l1_loss(recon_audio, target)
|
||||
|
||||
multi_stft_resolution_loss = 0.
|
||||
|
||||
for window_size in self.multi_stft_resolutions_window_sizes:
|
||||
res_stft_kwargs = dict(
|
||||
n_fft=max(window_size, self.multi_stft_n_fft), # not sure what n_fft is across multi resolution stft
|
||||
win_length=window_size,
|
||||
return_complex=True,
|
||||
window=self.multi_stft_window_fn(window_size, device=device),
|
||||
**self.multi_stft_kwargs,
|
||||
)
|
||||
|
||||
recon_Y = torch.stft(rearrange(recon_audio, '... s t -> (... s) t'), **res_stft_kwargs)
|
||||
target_Y = torch.stft(rearrange(target, '... s t -> (... s) t'), **res_stft_kwargs)
|
||||
|
||||
multi_stft_resolution_loss = multi_stft_resolution_loss + F.l1_loss(recon_Y, target_Y)
|
||||
|
||||
weighted_multi_resolution_loss = multi_stft_resolution_loss * self.multi_stft_resolution_loss_weight
|
||||
|
||||
total_loss = loss + weighted_multi_resolution_loss
|
||||
|
||||
if not return_loss_breakdown:
|
||||
return total_loss
|
||||
|
||||
return total_loss, (loss, multi_stft_resolution_loss)
|
||||
@@ -0,0 +1,703 @@
|
||||
from functools import partial
|
||||
|
||||
import torch
|
||||
from torch import nn, einsum, Tensor
|
||||
from torch.nn import Module, ModuleList
|
||||
import torch.nn.functional as F
|
||||
|
||||
from .attend import Attend
|
||||
try:
|
||||
from .attend_sage import Attend as AttendSage
|
||||
except:
|
||||
pass
|
||||
from torch.utils.checkpoint import checkpoint
|
||||
|
||||
from beartype.typing import Tuple, Optional, List, Callable
|
||||
from beartype import beartype
|
||||
|
||||
from rotary_embedding_torch import RotaryEmbedding
|
||||
|
||||
from einops import rearrange, pack, unpack, reduce, repeat
|
||||
from einops.layers.torch import Rearrange
|
||||
|
||||
from librosa import filters
|
||||
|
||||
|
||||
# helper functions
|
||||
|
||||
def exists(val):
|
||||
return val is not None
|
||||
|
||||
|
||||
def default(v, d):
|
||||
return v if exists(v) else d
|
||||
|
||||
|
||||
def pack_one(t, pattern):
|
||||
return pack([t], pattern)
|
||||
|
||||
|
||||
def unpack_one(t, ps, pattern):
|
||||
return unpack(t, ps, pattern)[0]
|
||||
|
||||
|
||||
def pad_at_dim(t, pad, dim=-1, value=0.):
|
||||
dims_from_right = (- dim - 1) if dim < 0 else (t.ndim - dim - 1)
|
||||
zeros = ((0, 0) * dims_from_right)
|
||||
return F.pad(t, (*zeros, *pad), value=value)
|
||||
|
||||
|
||||
def l2norm(t):
|
||||
return F.normalize(t, dim=-1, p=2)
|
||||
|
||||
|
||||
# norm
|
||||
|
||||
class RMSNorm(Module):
|
||||
def __init__(self, dim):
|
||||
super().__init__()
|
||||
self.scale = dim ** 0.5
|
||||
self.gamma = nn.Parameter(torch.ones(dim))
|
||||
|
||||
def forward(self, x):
|
||||
return F.normalize(x, dim=-1) * self.scale * self.gamma
|
||||
|
||||
|
||||
# attention
|
||||
|
||||
class FeedForward(Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
mult=4,
|
||||
dropout=0.
|
||||
):
|
||||
super().__init__()
|
||||
dim_inner = int(dim * mult)
|
||||
self.net = nn.Sequential(
|
||||
RMSNorm(dim),
|
||||
nn.Linear(dim, dim_inner),
|
||||
nn.GELU(),
|
||||
nn.Dropout(dropout),
|
||||
nn.Linear(dim_inner, dim),
|
||||
nn.Dropout(dropout)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
return self.net(x)
|
||||
|
||||
|
||||
class Attention(Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
heads=8,
|
||||
dim_head=64,
|
||||
dropout=0.,
|
||||
rotary_embed=None,
|
||||
flash=True,
|
||||
sage_attention=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.heads = heads
|
||||
self.scale = dim_head ** -0.5
|
||||
dim_inner = heads * dim_head
|
||||
|
||||
self.rotary_embed = rotary_embed
|
||||
|
||||
if sage_attention:
|
||||
self.attend = AttendSage(flash=flash, dropout=dropout)
|
||||
else:
|
||||
self.attend = Attend(flash=flash, dropout=dropout)
|
||||
self.norm = RMSNorm(dim)
|
||||
self.to_qkv = nn.Linear(dim, dim_inner * 3, bias=False)
|
||||
|
||||
self.to_gates = nn.Linear(dim, heads)
|
||||
|
||||
self.to_out = nn.Sequential(
|
||||
nn.Linear(dim_inner, dim, bias=False),
|
||||
nn.Dropout(dropout)
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = self.norm(x)
|
||||
|
||||
q, k, v = rearrange(self.to_qkv(x), 'b n (qkv h d) -> qkv b h n d', qkv=3, h=self.heads)
|
||||
|
||||
if exists(self.rotary_embed):
|
||||
q = self.rotary_embed.rotate_queries_or_keys(q)
|
||||
k = self.rotary_embed.rotate_queries_or_keys(k)
|
||||
|
||||
out = self.attend(q, k, v)
|
||||
|
||||
gates = self.to_gates(x)
|
||||
out = out * rearrange(gates, 'b n h -> b h n 1').sigmoid()
|
||||
|
||||
out = rearrange(out, 'b h n d -> b n (h d)')
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class LinearAttention(Module):
|
||||
"""
|
||||
this flavor of linear attention proposed in https://arxiv.org/abs/2106.09681 by El-Nouby et al.
|
||||
"""
|
||||
|
||||
@beartype
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
dim,
|
||||
dim_head=32,
|
||||
heads=8,
|
||||
scale=8,
|
||||
flash=False,
|
||||
dropout=0.,
|
||||
sage_attention=False
|
||||
):
|
||||
super().__init__()
|
||||
dim_inner = dim_head * heads
|
||||
self.norm = RMSNorm(dim)
|
||||
|
||||
self.to_qkv = nn.Sequential(
|
||||
nn.Linear(dim, dim_inner * 3, bias=False),
|
||||
Rearrange('b n (qkv h d) -> qkv b h d n', qkv=3, h=heads)
|
||||
)
|
||||
|
||||
self.temperature = nn.Parameter(torch.ones(heads, 1, 1))
|
||||
|
||||
if sage_attention:
|
||||
self.attend = AttendSage(
|
||||
scale=scale,
|
||||
dropout=dropout,
|
||||
flash=flash
|
||||
)
|
||||
else:
|
||||
self.attend = Attend(
|
||||
scale=scale,
|
||||
dropout=dropout,
|
||||
flash=flash
|
||||
)
|
||||
|
||||
self.to_out = nn.Sequential(
|
||||
Rearrange('b h d n -> b n (h d)'),
|
||||
nn.Linear(dim_inner, dim, bias=False)
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x
|
||||
):
|
||||
x = self.norm(x)
|
||||
|
||||
q, k, v = self.to_qkv(x)
|
||||
|
||||
q, k = map(l2norm, (q, k))
|
||||
q = q * self.temperature.exp()
|
||||
|
||||
out = self.attend(q, k, v)
|
||||
|
||||
return self.to_out(out)
|
||||
|
||||
|
||||
class Transformer(Module):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
dim,
|
||||
depth,
|
||||
dim_head=64,
|
||||
heads=8,
|
||||
attn_dropout=0.,
|
||||
ff_dropout=0.,
|
||||
ff_mult=4,
|
||||
norm_output=True,
|
||||
rotary_embed=None,
|
||||
flash_attn=True,
|
||||
linear_attn=False,
|
||||
sage_attention=False,
|
||||
):
|
||||
super().__init__()
|
||||
self.layers = ModuleList([])
|
||||
|
||||
for _ in range(depth):
|
||||
if linear_attn:
|
||||
attn = LinearAttention(
|
||||
dim=dim,
|
||||
dim_head=dim_head,
|
||||
heads=heads,
|
||||
dropout=attn_dropout,
|
||||
flash=flash_attn,
|
||||
sage_attention=sage_attention
|
||||
)
|
||||
else:
|
||||
attn = Attention(
|
||||
dim=dim,
|
||||
dim_head=dim_head,
|
||||
heads=heads,
|
||||
dropout=attn_dropout,
|
||||
rotary_embed=rotary_embed,
|
||||
flash=flash_attn,
|
||||
sage_attention=sage_attention
|
||||
)
|
||||
|
||||
self.layers.append(ModuleList([
|
||||
attn,
|
||||
FeedForward(dim=dim, mult=ff_mult, dropout=ff_dropout)
|
||||
]))
|
||||
|
||||
self.norm = RMSNorm(dim) if norm_output else nn.Identity()
|
||||
|
||||
def forward(self, x):
|
||||
|
||||
for attn, ff in self.layers:
|
||||
x = attn(x) + x
|
||||
x = ff(x) + x
|
||||
|
||||
return self.norm(x)
|
||||
|
||||
|
||||
# bandsplit module
|
||||
|
||||
class BandSplit(Module):
|
||||
@beartype
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
dim_inputs: Tuple[int, ...]
|
||||
):
|
||||
super().__init__()
|
||||
self.dim_inputs = dim_inputs
|
||||
self.to_features = ModuleList([])
|
||||
|
||||
for dim_in in dim_inputs:
|
||||
net = nn.Sequential(
|
||||
RMSNorm(dim_in),
|
||||
nn.Linear(dim_in, dim)
|
||||
)
|
||||
|
||||
self.to_features.append(net)
|
||||
|
||||
def forward(self, x):
|
||||
x = x.split(self.dim_inputs, dim=-1)
|
||||
|
||||
outs = []
|
||||
for split_input, to_feature in zip(x, self.to_features):
|
||||
split_output = to_feature(split_input)
|
||||
outs.append(split_output)
|
||||
|
||||
return torch.stack(outs, dim=-2)
|
||||
|
||||
|
||||
def MLP(
|
||||
dim_in,
|
||||
dim_out,
|
||||
dim_hidden=None,
|
||||
depth=1,
|
||||
activation=nn.Tanh
|
||||
):
|
||||
dim_hidden = default(dim_hidden, dim_in)
|
||||
|
||||
net = []
|
||||
dims = (dim_in, *((dim_hidden,) * depth), dim_out)
|
||||
|
||||
for ind, (layer_dim_in, layer_dim_out) in enumerate(zip(dims[:-1], dims[1:])):
|
||||
is_last = ind == (len(dims) - 2)
|
||||
|
||||
net.append(nn.Linear(layer_dim_in, layer_dim_out))
|
||||
|
||||
if is_last:
|
||||
continue
|
||||
|
||||
net.append(activation())
|
||||
|
||||
return nn.Sequential(*net)
|
||||
|
||||
|
||||
class MaskEstimator(Module):
|
||||
@beartype
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
dim_inputs: Tuple[int, ...],
|
||||
depth,
|
||||
mlp_expansion_factor=4
|
||||
):
|
||||
super().__init__()
|
||||
self.dim_inputs = dim_inputs
|
||||
self.to_freqs = ModuleList([])
|
||||
dim_hidden = dim * mlp_expansion_factor
|
||||
|
||||
for dim_in in dim_inputs:
|
||||
net = []
|
||||
|
||||
mlp = nn.Sequential(
|
||||
MLP(dim, dim_in * 2, dim_hidden=dim_hidden, depth=depth),
|
||||
nn.GLU(dim=-1)
|
||||
)
|
||||
|
||||
self.to_freqs.append(mlp)
|
||||
|
||||
def forward(self, x):
|
||||
x = x.unbind(dim=-2)
|
||||
|
||||
outs = []
|
||||
|
||||
for band_features, mlp in zip(x, self.to_freqs):
|
||||
freq_out = mlp(band_features)
|
||||
outs.append(freq_out)
|
||||
|
||||
return torch.cat(outs, dim=-1)
|
||||
|
||||
|
||||
# main class
|
||||
|
||||
class MelBandRoformer(Module):
|
||||
|
||||
@beartype
|
||||
def __init__(
|
||||
self,
|
||||
dim,
|
||||
*,
|
||||
depth,
|
||||
stereo=False,
|
||||
num_stems=1,
|
||||
time_transformer_depth=2,
|
||||
freq_transformer_depth=2,
|
||||
linear_transformer_depth=0,
|
||||
num_bands=60,
|
||||
dim_head=64,
|
||||
heads=8,
|
||||
attn_dropout=0.1,
|
||||
ff_dropout=0.1,
|
||||
flash_attn=True,
|
||||
dim_freqs_in=1025,
|
||||
sample_rate=44100, # needed for mel filter bank from librosa
|
||||
stft_n_fft=2048,
|
||||
stft_hop_length=512,
|
||||
# 10ms at 44100Hz, from sections 4.1, 4.4 in the paper - @faroit recommends // 2 or // 4 for better reconstruction
|
||||
stft_win_length=2048,
|
||||
stft_normalized=False,
|
||||
stft_window_fn: Optional[Callable] = None,
|
||||
mask_estimator_depth=1,
|
||||
multi_stft_resolution_loss_weight=1.,
|
||||
multi_stft_resolutions_window_sizes: Tuple[int, ...] = (4096, 2048, 1024, 512, 256),
|
||||
multi_stft_hop_size=147,
|
||||
multi_stft_normalized=False,
|
||||
multi_stft_window_fn: Callable = torch.hann_window,
|
||||
match_input_audio_length=False, # if True, pad output tensor to match length of input tensor
|
||||
mlp_expansion_factor=4,
|
||||
use_torch_checkpoint=False,
|
||||
skip_connection=False,
|
||||
sage_attention=False,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.stereo = stereo
|
||||
self.audio_channels = 2 if stereo else 1
|
||||
self.num_stems = num_stems
|
||||
self.use_torch_checkpoint = use_torch_checkpoint
|
||||
self.skip_connection = skip_connection
|
||||
|
||||
self.layers = ModuleList([])
|
||||
|
||||
if sage_attention:
|
||||
print("Use Sage Attention")
|
||||
|
||||
transformer_kwargs = dict(
|
||||
dim=dim,
|
||||
heads=heads,
|
||||
dim_head=dim_head,
|
||||
attn_dropout=attn_dropout,
|
||||
ff_dropout=ff_dropout,
|
||||
flash_attn=flash_attn,
|
||||
sage_attention=sage_attention,
|
||||
)
|
||||
|
||||
time_rotary_embed = RotaryEmbedding(dim=dim_head)
|
||||
freq_rotary_embed = RotaryEmbedding(dim=dim_head)
|
||||
|
||||
for _ in range(depth):
|
||||
tran_modules = []
|
||||
if linear_transformer_depth > 0:
|
||||
tran_modules.append(Transformer(depth=linear_transformer_depth, linear_attn=True, **transformer_kwargs))
|
||||
tran_modules.append(
|
||||
Transformer(depth=time_transformer_depth, rotary_embed=time_rotary_embed, **transformer_kwargs)
|
||||
)
|
||||
tran_modules.append(
|
||||
Transformer(depth=freq_transformer_depth, rotary_embed=freq_rotary_embed, **transformer_kwargs)
|
||||
)
|
||||
self.layers.append(nn.ModuleList(tran_modules))
|
||||
|
||||
self.stft_window_fn = partial(default(stft_window_fn, torch.hann_window), stft_win_length)
|
||||
|
||||
self.stft_kwargs = dict(
|
||||
n_fft=stft_n_fft,
|
||||
hop_length=stft_hop_length,
|
||||
win_length=stft_win_length,
|
||||
normalized=stft_normalized
|
||||
)
|
||||
|
||||
freqs = torch.stft(torch.randn(1, 4096), **self.stft_kwargs, window=torch.ones(stft_n_fft), return_complex=True).shape[1]
|
||||
|
||||
# create mel filter bank
|
||||
# with librosa.filters.mel as in section 2 of paper
|
||||
|
||||
mel_filter_bank_numpy = filters.mel(sr=sample_rate, n_fft=stft_n_fft, n_mels=num_bands)
|
||||
|
||||
mel_filter_bank = torch.from_numpy(mel_filter_bank_numpy)
|
||||
|
||||
# for some reason, it doesn't include the first freq? just force a value for now
|
||||
|
||||
mel_filter_bank[0][0] = 1.
|
||||
|
||||
# In some systems/envs we get 0.0 instead of ~1.9e-18 in the last position,
|
||||
# so let's force a positive value
|
||||
|
||||
mel_filter_bank[-1, -1] = 1.
|
||||
|
||||
# binary as in paper (then estimated masks are averaged for overlapping regions)
|
||||
|
||||
freqs_per_band = mel_filter_bank > 0
|
||||
assert freqs_per_band.any(dim=0).all(), 'all frequencies need to be covered by all bands for now'
|
||||
|
||||
repeated_freq_indices = repeat(torch.arange(freqs), 'f -> b f', b=num_bands)
|
||||
freq_indices = repeated_freq_indices[freqs_per_band]
|
||||
|
||||
if stereo:
|
||||
freq_indices = repeat(freq_indices, 'f -> f s', s=2)
|
||||
freq_indices = freq_indices * 2 + torch.arange(2)
|
||||
freq_indices = rearrange(freq_indices, 'f s -> (f s)')
|
||||
|
||||
self.register_buffer('freq_indices', freq_indices, persistent=False)
|
||||
self.register_buffer('freqs_per_band', freqs_per_band, persistent=False)
|
||||
|
||||
num_freqs_per_band = reduce(freqs_per_band, 'b f -> b', 'sum')
|
||||
num_bands_per_freq = reduce(freqs_per_band, 'b f -> f', 'sum')
|
||||
|
||||
self.register_buffer('num_freqs_per_band', num_freqs_per_band, persistent=False)
|
||||
self.register_buffer('num_bands_per_freq', num_bands_per_freq, persistent=False)
|
||||
|
||||
# band split and mask estimator
|
||||
|
||||
freqs_per_bands_with_complex = tuple(2 * f * self.audio_channels for f in num_freqs_per_band.tolist())
|
||||
|
||||
self.band_split = BandSplit(
|
||||
dim=dim,
|
||||
dim_inputs=freqs_per_bands_with_complex
|
||||
)
|
||||
|
||||
self.mask_estimators = nn.ModuleList([])
|
||||
|
||||
for _ in range(num_stems):
|
||||
mask_estimator = MaskEstimator(
|
||||
dim=dim,
|
||||
dim_inputs=freqs_per_bands_with_complex,
|
||||
depth=mask_estimator_depth,
|
||||
mlp_expansion_factor=mlp_expansion_factor,
|
||||
)
|
||||
|
||||
self.mask_estimators.append(mask_estimator)
|
||||
|
||||
# for the multi-resolution stft loss
|
||||
|
||||
self.multi_stft_resolution_loss_weight = multi_stft_resolution_loss_weight
|
||||
self.multi_stft_resolutions_window_sizes = multi_stft_resolutions_window_sizes
|
||||
self.multi_stft_n_fft = stft_n_fft
|
||||
self.multi_stft_window_fn = multi_stft_window_fn
|
||||
|
||||
self.multi_stft_kwargs = dict(
|
||||
hop_length=multi_stft_hop_size,
|
||||
normalized=multi_stft_normalized
|
||||
)
|
||||
|
||||
self.match_input_audio_length = match_input_audio_length
|
||||
|
||||
def forward(
|
||||
self,
|
||||
raw_audio,
|
||||
target=None,
|
||||
return_loss_breakdown=False
|
||||
):
|
||||
"""
|
||||
einops
|
||||
|
||||
b - batch
|
||||
f - freq
|
||||
t - time
|
||||
s - audio channel (1 for mono, 2 for stereo)
|
||||
n - number of 'stems'
|
||||
c - complex (2)
|
||||
d - feature dimension
|
||||
"""
|
||||
|
||||
device = raw_audio.device
|
||||
|
||||
if raw_audio.ndim == 2:
|
||||
raw_audio = rearrange(raw_audio, 'b t -> b 1 t')
|
||||
|
||||
batch, channels, raw_audio_length = raw_audio.shape
|
||||
|
||||
istft_length = raw_audio_length if self.match_input_audio_length else None
|
||||
|
||||
assert (not self.stereo and channels == 1) or (
|
||||
self.stereo and channels == 2), 'stereo needs to be set to True if passing in audio signal that is stereo (channel dimension of 2). also need to be False if mono (channel dimension of 1)'
|
||||
|
||||
# to stft
|
||||
|
||||
raw_audio, batch_audio_channel_packed_shape = pack_one(raw_audio, '* t')
|
||||
|
||||
stft_window = self.stft_window_fn(device=device)
|
||||
|
||||
stft_repr = torch.stft(raw_audio, **self.stft_kwargs, window=stft_window, return_complex=True)
|
||||
stft_repr = torch.view_as_real(stft_repr)
|
||||
|
||||
stft_repr = unpack_one(stft_repr, batch_audio_channel_packed_shape, '* f t c')
|
||||
|
||||
# merge stereo / mono into the frequency, with frequency leading dimension, for band splitting
|
||||
stft_repr = rearrange(stft_repr,'b s f t c -> b (f s) t c')
|
||||
|
||||
# index out all frequencies for all frequency ranges across bands ascending in one go
|
||||
|
||||
batch_arange = torch.arange(batch, device=device)[..., None]
|
||||
|
||||
# account for stereo
|
||||
|
||||
x = stft_repr[batch_arange, self.freq_indices]
|
||||
|
||||
# fold the complex (real and imag) into the frequencies dimension
|
||||
|
||||
x = rearrange(x, 'b f t c -> b t (f c)')
|
||||
|
||||
if self.use_torch_checkpoint:
|
||||
x = checkpoint(self.band_split, x, use_reentrant=False)
|
||||
else:
|
||||
x = self.band_split(x)
|
||||
|
||||
# axial / hierarchical attention
|
||||
|
||||
store = [None] * len(self.layers)
|
||||
for i, transformer_block in enumerate(self.layers):
|
||||
|
||||
if len(transformer_block) == 3:
|
||||
linear_transformer, time_transformer, freq_transformer = transformer_block
|
||||
|
||||
x, ft_ps = pack([x], 'b * d')
|
||||
if self.use_torch_checkpoint:
|
||||
x = checkpoint(linear_transformer, x, use_reentrant=False)
|
||||
else:
|
||||
x = linear_transformer(x)
|
||||
x, = unpack(x, ft_ps, 'b * d')
|
||||
else:
|
||||
time_transformer, freq_transformer = transformer_block
|
||||
|
||||
if self.skip_connection:
|
||||
# Sum all previous
|
||||
for j in range(i):
|
||||
x = x + store[j]
|
||||
|
||||
x = rearrange(x, 'b t f d -> b f t d')
|
||||
x, ps = pack([x], '* t d')
|
||||
|
||||
if self.use_torch_checkpoint:
|
||||
x = checkpoint(time_transformer, x, use_reentrant=False)
|
||||
else:
|
||||
x = time_transformer(x)
|
||||
|
||||
x, = unpack(x, ps, '* t d')
|
||||
x = rearrange(x, 'b f t d -> b t f d')
|
||||
x, ps = pack([x], '* f d')
|
||||
|
||||
if self.use_torch_checkpoint:
|
||||
x = checkpoint(freq_transformer, x, use_reentrant=False)
|
||||
else:
|
||||
x = freq_transformer(x)
|
||||
|
||||
x, = unpack(x, ps, '* f d')
|
||||
|
||||
if self.skip_connection:
|
||||
store[i] = x
|
||||
|
||||
num_stems = len(self.mask_estimators)
|
||||
if self.use_torch_checkpoint:
|
||||
masks = torch.stack([checkpoint(fn, x, use_reentrant=False) for fn in self.mask_estimators], dim=1)
|
||||
else:
|
||||
masks = torch.stack([fn(x) for fn in self.mask_estimators], dim=1)
|
||||
masks = rearrange(masks, 'b n t (f c) -> b n f t c', c=2)
|
||||
|
||||
# modulate frequency representation
|
||||
|
||||
stft_repr = rearrange(stft_repr, 'b f t c -> b 1 f t c')
|
||||
|
||||
# complex number multiplication
|
||||
|
||||
stft_repr = torch.view_as_complex(stft_repr)
|
||||
masks = torch.view_as_complex(masks)
|
||||
|
||||
masks = masks.type(stft_repr.dtype)
|
||||
|
||||
# need to average the estimated mask for the overlapped frequencies
|
||||
|
||||
scatter_indices = repeat(self.freq_indices, 'f -> b n f t', b=batch, n=num_stems, t=stft_repr.shape[-1])
|
||||
|
||||
stft_repr_expanded_stems = repeat(stft_repr, 'b 1 ... -> b n ...', n=num_stems)
|
||||
masks_summed = torch.zeros_like(stft_repr_expanded_stems).scatter_add_(2, scatter_indices, masks)
|
||||
|
||||
denom = repeat(self.num_bands_per_freq, 'f -> (f r) 1', r=channels)
|
||||
|
||||
masks_averaged = masks_summed / denom.clamp(min=1e-8)
|
||||
|
||||
# modulate stft repr with estimated mask
|
||||
|
||||
stft_repr = stft_repr * masks_averaged
|
||||
|
||||
# istft
|
||||
|
||||
stft_repr = rearrange(stft_repr, 'b n (f s) t -> (b n s) f t', s=self.audio_channels)
|
||||
|
||||
recon_audio = torch.istft(stft_repr, **self.stft_kwargs, window=stft_window, return_complex=False,
|
||||
length=istft_length)
|
||||
|
||||
recon_audio = rearrange(recon_audio, '(b n s) t -> b n s t', b=batch, s=self.audio_channels, n=num_stems)
|
||||
|
||||
if num_stems == 1:
|
||||
recon_audio = rearrange(recon_audio, 'b 1 s t -> b s t')
|
||||
|
||||
# if a target is passed in, calculate loss for learning
|
||||
|
||||
if not exists(target):
|
||||
return recon_audio
|
||||
|
||||
if self.num_stems > 1:
|
||||
assert target.ndim == 4 and target.shape[1] == self.num_stems
|
||||
|
||||
if target.ndim == 2:
|
||||
target = rearrange(target, '... t -> ... 1 t')
|
||||
|
||||
target = target[..., :recon_audio.shape[-1]] # protect against lost length on istft
|
||||
|
||||
loss = F.l1_loss(recon_audio, target)
|
||||
|
||||
multi_stft_resolution_loss = 0.
|
||||
|
||||
for window_size in self.multi_stft_resolutions_window_sizes:
|
||||
res_stft_kwargs = dict(
|
||||
n_fft=max(window_size, self.multi_stft_n_fft), # not sure what n_fft is across multi resolution stft
|
||||
win_length=window_size,
|
||||
return_complex=True,
|
||||
window=self.multi_stft_window_fn(window_size, device=device),
|
||||
**self.multi_stft_kwargs,
|
||||
)
|
||||
|
||||
recon_Y = torch.stft(rearrange(recon_audio, '... s t -> (... s) t'), **res_stft_kwargs)
|
||||
target_Y = torch.stft(rearrange(target, '... s t -> (... s) t'), **res_stft_kwargs)
|
||||
|
||||
multi_stft_resolution_loss = multi_stft_resolution_loss + F.l1_loss(recon_Y, target_Y)
|
||||
|
||||
weighted_multi_resolution_loss = multi_stft_resolution_loss * self.multi_stft_resolution_loss_weight
|
||||
|
||||
total_loss = loss + weighted_multi_resolution_loss
|
||||
|
||||
if not return_loss_breakdown:
|
||||
return total_loss
|
||||
|
||||
return total_loss, (loss, multi_stft_resolution_loss)
|
||||
@@ -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