ok
This commit is contained in:
31
tests/test_matcha_flow.py
Normal file
31
tests/test_matcha_flow.py
Normal file
@@ -0,0 +1,31 @@
|
||||
import torch
|
||||
|
||||
from style_bert_vits2.models.matcha_flow import MatchaFlow
|
||||
|
||||
|
||||
def test_matcha_flow_loss_and_sampling_with_odd_length() -> None:
|
||||
model = MatchaFlow(
|
||||
latent_channels=8,
|
||||
speaker_channels=4,
|
||||
channels=16,
|
||||
num_heads=2,
|
||||
dropout=0.0,
|
||||
)
|
||||
target = torch.randn(2, 8, 7)
|
||||
mu = torch.randn_like(target)
|
||||
mask = torch.tensor(
|
||||
[[[1, 1, 1, 1, 1, 1, 1]], [[1, 1, 1, 1, 1, 0, 0]]],
|
||||
dtype=target.dtype,
|
||||
)
|
||||
speaker = torch.randn(2, 4, 1)
|
||||
|
||||
loss = model.compute_loss(target, mask, mu, speaker)
|
||||
loss.backward()
|
||||
|
||||
assert loss.ndim == 0
|
||||
assert torch.isfinite(loss)
|
||||
assert any(parameter.grad is not None for parameter in model.parameters())
|
||||
|
||||
sample = model(mu, mask, n_timesteps=2, speaker=speaker)
|
||||
assert sample.shape == target.shape
|
||||
assert torch.count_nonzero(sample[1, :, 5:]) == 0
|
||||
Reference in New Issue
Block a user