MCPcopy Create free account
hub / github.com/FireRedTeam/FireRedASR / Adapter

Class Adapter

fireredasr/models/module/adapter.py:5–30  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

3
4
5class Adapter(nn.Module):
6 def __init__(self, encoder_dim, llm_dim, downsample_rate=2):
7 super().__init__()
8 self.ds = downsample_rate
9 self.linear1 = nn.Linear(encoder_dim * downsample_rate, llm_dim)
10 self.relu = nn.ReLU()
11 self.linear2 = nn.Linear(llm_dim, llm_dim)
12
13 def forward(self, x, x_lens):
14 batch_size, seq_len, feat_dim = x.size()
15 num_frames_to_discard = seq_len % self.ds
16 if num_frames_to_discard > 0:
17 x = x[:, :-num_frames_to_discard, :]
18 seq_len = x.size(1)
19
20 x = x.contiguous()
21 x = x.view(
22 batch_size, seq_len // self.ds, feat_dim * self.ds
23 )
24
25 x = self.linear1(x)
26 x = self.relu(x)
27 x = self.linear2(x)
28
29 new_x_lens = torch.clamp(x_lens, max=seq_len) // self.ds
30 return x, new_x_lens

Callers 1

from_argsMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected