MCPcopy Create free account
hub / github.com/FireRedTeam/FireRedTTS2 / make_pad_mask

Function make_pad_mask

fireredtts2/codec/utils.py:5–12  ·  view source on GitHub ↗
(lengths: torch.Tensor, max_len: int = 0)

Source from the content-addressed store, hash-verified

3
4
5def make_pad_mask(lengths: torch.Tensor, max_len: int = 0) -> torch.Tensor:
6 batch_size = lengths.size(0)
7 max_len = max_len if max_len > 0 else lengths.max().item()
8 seq_range = torch.arange(0, max_len, dtype=torch.int64, device=lengths.device)
9 seq_range_expand = seq_range.unsqueeze(0).expand(batch_size, max_len)
10 seq_length_expand = lengths.unsqueeze(-1)
11 mask = seq_range_expand >= seq_length_expand
12 return mask # (b, t)
13
14
15def make_nonpad_mask(lengths: torch.Tensor, max_len: int = 0) -> torch.Tensor:

Callers 1

make_nonpad_maskFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected