95 lines
3.0 KiB
Python
95 lines
3.0 KiB
Python
# Copyright (c) Meta Platforms, Inc. and affiliates.
|
|
# All rights reserved.
|
|
#
|
|
# This source code is licensed under the license found in the
|
|
# LICENSE file in the root directory of this source tree.
|
|
# License: MIT
|
|
# Misc stuff used by Demucs
|
|
import functools
|
|
import math
|
|
import torch
|
|
from torch.nn import functional as F
|
|
import typing as tp
|
|
|
|
|
|
# ############################################################################################################################
|
|
# utils.py
|
|
# ############################################################################################################################
|
|
|
|
|
|
def unfold(a, kernel_size, stride):
|
|
"""Given input of size [*X, T], output Tensor of size [*X, F, K]
|
|
with K the kernel size, by extracting frames with the given stride.
|
|
|
|
This will pad the input so that `F = ceil(T / K)`.
|
|
|
|
see https://github.com/pytorch/pytorch/issues/60466
|
|
"""
|
|
*shape, length = a.shape
|
|
n_frames = math.ceil(length / stride)
|
|
tgt_length = (n_frames - 1) * stride + kernel_size
|
|
a = F.pad(a, (0, tgt_length - length))
|
|
strides = list(a.stride())
|
|
assert strides[-1] == 1, 'data should be contiguous'
|
|
strides = strides[:-1] + [stride, 1]
|
|
return a.as_strided([*shape, n_frames, kernel_size], strides)
|
|
|
|
|
|
def center_trim(tensor: torch.Tensor, reference: tp.Union[torch.Tensor, int]):
|
|
"""
|
|
Center trim `tensor` with respect to `reference`, along the last dimension.
|
|
`reference` can also be a number, representing the length to trim to.
|
|
If the size difference != 0 mod 2, the extra sample is removed on the right side.
|
|
"""
|
|
ref_size: int
|
|
if isinstance(reference, torch.Tensor):
|
|
ref_size = reference.size(-1)
|
|
else:
|
|
ref_size = reference
|
|
delta = tensor.size(-1) - ref_size
|
|
if delta < 0:
|
|
raise ValueError("tensor must be larger than reference. " f"Delta is {delta}.")
|
|
if delta:
|
|
tensor = tensor[..., delta // 2:-(delta - delta // 2)]
|
|
return tensor
|
|
|
|
|
|
class DummyPoolExecutor:
|
|
class DummyResult:
|
|
def __init__(self, func, *args, **kwargs):
|
|
self.func = func
|
|
self.args = args
|
|
self.kwargs = kwargs
|
|
|
|
def result(self):
|
|
return self.func(*self.args, **self.kwargs)
|
|
|
|
def __init__(self, workers=0):
|
|
pass
|
|
|
|
def submit(self, func, *args, **kwargs):
|
|
return DummyPoolExecutor.DummyResult(func, *args, **kwargs)
|
|
|
|
def shutdown(self, wait=True, cancel_futures=True):
|
|
return
|
|
|
|
def __enter__(self):
|
|
return self
|
|
|
|
def __exit__(self, exc_type, exc_value, exc_tb):
|
|
return
|
|
|
|
|
|
# ############################################################################################################################
|
|
# state.py
|
|
# ############################################################################################################################
|
|
|
|
|
|
def capture_init(init):
|
|
@functools.wraps(init)
|
|
def __init__(self, *args, **kwargs):
|
|
self._init_args_kwargs = (args, kwargs)
|
|
init(self, *args, **kwargs)
|
|
|
|
return __init__
|