Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions configs/acoustic.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,7 @@ max_beta: 0.02
enc_ffn_kernel_size: 3
use_rope: true
rope_interleaved: false
rope_theta: 10000
use_stretch_embed: true
use_variance_scaling: true
rel_pos: true
Expand Down
1 change: 1 addition & 0 deletions configs/templates/config_acoustic.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -75,6 +75,7 @@ diffusion_type: reflow
enc_ffn_kernel_size: 3
use_rope: true
rope_interleaved: false
rope_theta: 10000
use_stretch_embed: true
use_variance_scaling: true
use_shallow_diffusion: true
Expand Down
1 change: 1 addition & 0 deletions configs/templates/config_variance.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,7 @@ tension_logit_max: 10.0
enc_ffn_kernel_size: 3
use_rope: true
rope_interleaved: false
rope_theta: 10000
use_stretch_embed: false
use_variance_scaling: true
hidden_size: 384
Expand Down
1 change: 1 addition & 0 deletions configs/variance.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,7 @@ predict_tension: false
enc_ffn_kernel_size: 3
use_rope: true
rope_interleaved: false
rope_theta: 10000
use_stretch_embed: false
use_variance_scaling: true
rel_pos: true
Expand Down
54 changes: 26 additions & 28 deletions modules/commons/rotary_embedding_torch.py
Original file line number Diff line number Diff line change
@@ -1,29 +1,22 @@
import torch
from einops import rearrange, repeat
from torch import einsum, Tensor
from torch import Tensor
from torch.nn import Module


def rotate_half(x: Tensor, interleaved=True) -> Tensor:
if not interleaved:
# x_half1, x_half2 = x.chunk(2, dim=-1)
# Using torch.split instead of chunk for ONNX export compatibility.
x1, x2 = torch.split(x, x.size(-1) // 2, dim=-1)
return torch.cat((-x2, x1), dim=-1)
else:
x = rearrange(x, '... (d r) -> ... d r', r=2)
x1, x2 = x.unbind(dim=-1)
x = torch.stack((-x2, x1), dim=-1)
return rearrange(x, '... d r -> ... (d r)')


def apply_rotary_emb(freqs: Tensor, t: Tensor, interleaved=True) -> Tensor:
rot_dim = freqs.shape[-1]
def apply_rotary_emb(freqs_cos: Tensor, freqs_sin: Tensor, t: Tensor, interleaved=True) -> Tensor:
rot_dim = freqs_cos.shape[-1]
t_to_rotate = t[..., :rot_dim]
t_pass_through = t[..., rot_dim:]

t_rotated = (t_to_rotate * freqs.cos()) + (rotate_half(t_to_rotate, interleaved) * freqs.sin())
if interleaved:
x = t_to_rotate.view(*t_to_rotate.shape[:-1], t_to_rotate.size(-1) // 2, 2)
x1, x2 = x.unbind(dim=-1)
rotated_half = torch.stack((-x2, x1), dim=-1).reshape_as(t_to_rotate)
else:
x1, x2 = torch.split(t_to_rotate, t_to_rotate.size(-1) // 2, dim=-1)
rotated_half = torch.cat((-x2, x1), dim=-1)

t_rotated = (t_to_rotate * freqs_cos) + (rotated_half * freqs_sin)
return torch.cat((t_rotated, t_pass_through), dim=-1)


Expand All @@ -40,24 +33,29 @@ def __init__(
self.cached_freqs_seq_len = max_seq_len
inv_freq = 1. / (theta ** (torch.arange(0, dim, 2).float() / dim))
self.register_buffer('inv_freq', inv_freq, persistent=False)
self.register_buffer('cached_freqs', self._precompute_cache(max_seq_len), persistent=False)
cos, sin = self._precompute_cache(max_seq_len)
self.register_buffer('cached_cos', cos, persistent=False)
self.register_buffer('cached_sin', sin, persistent=False)

def _precompute_cache(self, seq_len: int):
seq = torch.arange(seq_len, device=self.inv_freq.device, dtype=self.inv_freq.dtype)
freqs = einsum('i, j -> i j', seq, self.inv_freq)
# Cache fp32 cos/sin, cast only at use — fp16/bf16 training must not
# recompute trig on low-precision angles.
seq = torch.arange(seq_len, device=self.inv_freq.device, dtype=torch.float32)
freqs = torch.einsum('i, j -> i j', seq, self.inv_freq.float())
if self.interleaved:
freqs = repeat(freqs, '... n -> ... (n r)', r=2)
freqs = torch.repeat_interleave(freqs, 2, dim=-1)
else:
freqs = torch.cat((freqs, freqs), dim=-1)
return freqs
return torch.cos(freqs), torch.sin(freqs)

def forward(self, seq_len: int) -> Tensor:
def forward(self, seq_len: int):
if seq_len > self.cached_freqs_seq_len:
raise RuntimeError("sequence exceeds RoPE max_seq_len!")
return self.cached_freqs[0: seq_len].detach()
return self.cached_cos[0: seq_len].detach(), self.cached_sin[0: seq_len].detach()

def rotate_queries_or_keys(self, t: Tensor) -> Tensor:
device, dtype, seq_len = t.device, t.dtype, t.shape[-2]
freqs = self.forward(seq_len=seq_len)

return apply_rotary_emb(freqs.to(device=device, dtype=dtype), t, self.interleaved)
freqs_cos, freqs_sin = self.forward(seq_len=seq_len)
freqs_cos = freqs_cos.to(device=device, dtype=dtype)
freqs_sin = freqs_sin.to(device=device, dtype=dtype)
return apply_rotary_emb(freqs_cos, freqs_sin, t, self.interleaved)
1 change: 1 addition & 0 deletions modules/fastspeech/acoustic_encoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,7 @@ def __init__(self, vocab_size):
dropout=hparams['dropout'], num_heads=hparams['num_heads'],
use_pos_embed=hparams['use_pos_embed'], rel_pos=hparams.get('rel_pos', False),
use_rope=hparams.get('use_rope', False), rope_interleaved=hparams.get('rope_interleaved', True),
rope_theta=hparams.get('rope_theta', 10000),
mix_ln_layer=self.mix_ln_layer
)

Expand Down
6 changes: 4 additions & 2 deletions modules/fastspeech/tts_modules.py
Original file line number Diff line number Diff line change
Expand Up @@ -373,7 +373,7 @@ def __init__(
self, hidden_size, num_layers,
ffn_kernel_size=9, ffn_act='gelu',
dropout=None, num_heads=2, use_pos_embed=True, rel_pos=True,
use_rope=False, rope_interleaved=True, mix_ln_layer=None
use_rope=False, rope_interleaved=True, rope_theta=10000, mix_ln_layer=None
):
super().__init__()
self.num_layers = num_layers
Expand All @@ -386,7 +386,9 @@ def __init__(
"RoPE requires the hidden size to be multiple of "
f"num_heads * 2 = {num_heads * 2}, but got {embed_dim}."
)
rotary_embed = RotaryEmbedding(dim=embed_dim // num_heads, interleaved=rope_interleaved)
rotary_embed = RotaryEmbedding(
dim=embed_dim // num_heads, theta=rope_theta, interleaved=rope_interleaved
)
else:
rotary_embed = None
self.layers = nn.ModuleList([
Expand Down
6 changes: 4 additions & 2 deletions modules/fastspeech/variance_encoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,8 @@ def __init__(self, vocab_size):
ffn_kernel_size=hparams['enc_ffn_kernel_size'], ffn_act=hparams['ffn_act'],
dropout=hparams['dropout'], num_heads=hparams['num_heads'],
use_pos_embed=hparams['use_pos_embed'], rel_pos=hparams.get('rel_pos', False),
use_rope=hparams.get('use_rope', False), rope_interleaved=hparams.get('rope_interleaved', True)
use_rope=hparams.get('use_rope', False), rope_interleaved=hparams.get('rope_interleaved', True),
rope_theta=hparams.get('rope_theta', 10000)
)

dur_hparams = hparams['dur_prediction_args']
Expand Down Expand Up @@ -128,7 +129,8 @@ def get_hparam(key):
ffn_kernel_size=get_hparam('enc_ffn_kernel_size'), ffn_act=get_hparam('ffn_act'),
dropout=get_hparam('dropout'), num_heads=get_hparam('num_heads'),
use_pos_embed=get_hparam('use_pos_embed'), rel_pos=get_hparam('rel_pos'),
use_rope=get_hparam('use_rope'), rope_interleaved=hparams.get('rope_interleaved', True)
use_rope=get_hparam('use_rope'), rope_interleaved=hparams.get('rope_interleaved', True),
rope_theta=hparams.get('rope_theta', 10000)
)
self.out_proj = Linear(hidden_size, hparams['hidden_size'])

Expand Down
1 change: 0 additions & 1 deletion requirements.txt
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,6 @@
# See instructions at https://pytorch.org/get-started/locally/

click
einops>=0.7.0
h5py
librosa<0.10.0
lightning~=2.3.0
Expand Down