From 378eedbdc348a92774dd3a6183a4ec5d6b977dde Mon Sep 17 00:00:00 2001
From: 游雁 <zhifu.gzf@alibaba-inc.com>
Date: 星期三, 05 六月 2024 09:31:56 +0800
Subject: [PATCH] auto frontend

---
 funasr/models/transformer/encoder.py |    2 
 funasr/models/llm_asr/adaptor.py     |   62 +++++++++++++++++++++++++++++++
 2 files changed, 63 insertions(+), 1 deletions(-)

diff --git a/funasr/models/llm_asr/adaptor.py b/funasr/models/llm_asr/adaptor.py
index 8c2a804..75494de 100644
--- a/funasr/models/llm_asr/adaptor.py
+++ b/funasr/models/llm_asr/adaptor.py
@@ -63,3 +63,65 @@
         query_proj = self.norm(self.linear(query_output.last_hidden_state))
 
         return query_proj
+
+
+@tables.register("adaptor_classes", "Transformer")
+class Transformer(nn.Module):
+    def __init__(
+        self, downsample_rate=2, encoder_dim=1280, llm_dim=4096, ffn_dim: int = 2048, **kwargs
+    ):
+        super().__init__()
+        self.k = downsample_rate
+        self.encoder_dim = encoder_dim
+        self.llm_dim = llm_dim
+        self.linear1 = nn.Linear(self.encoder_dim * self.k, ffn_dim)
+        self.relu = nn.ReLU()
+        self.linear2 = nn.Linear(ffn_dim, self.llm_dim)
+        from funasr.models.transformer.encoder import EncoderLayer
+        from funasr.models.transformer.attention import MultiHeadedAttention
+        from funasr.models.transformer.positionwise_feed_forward import PositionwiseFeedForward
+
+        self.blocks = nn.ModuleList(
+            [
+                EncoderLayer(
+                    output_size,
+                    MultiHeadedAttention(
+                        kwargs.get("attention_heads", 8),
+                        llm_dim,
+                        kwargs.get("attention_dropout_rate", 0.0),
+                    ),
+                    positionwise_layer(
+                        llm_dim,
+                        llm_dim // 4,
+                        kwargs.get("dropout_rate", 0.0),
+                    ),
+                    kwargs.get("dropout_rate", 0.0),
+                )
+                for i in range(kwargs.get("n_layer", 2))
+            ]
+        )
+
+    def forward(self, x, ilens=None):
+
+        batch_size, seq_len, dim = x.size()
+        # num_frames_to_discard = seq_len % self.k
+        chunk_num = (seq_len - 1) // self.k + 1
+        pad_num = chunk_num * self.k - seq_len
+        x = F.pad(x, (0, 0, 0, pad_num, 0, 0), value=0.0)
+        # if num_frames_to_discard > 0:
+        #     x = x[:, :-num_frames_to_discard, :]
+        seq_len = x.size(1)
+
+        x = x.contiguous()
+        x = x.view(batch_size, chunk_num, dim * self.k)
+        x = self.linear1(x)
+        x = self.relu(x)
+        x = self.linear2(x)
+
+        olens = None
+        if ilens is not None:
+            olens = (ilens - 1) // self.k + 1
+            mask = (~make_pad_mask(olens)[:, None, :]).to(x.device)
+        for layer, block in enumerate(self.blocks):
+            x, masks = block(x, masks)
+        return x, olens
diff --git a/funasr/models/transformer/encoder.py b/funasr/models/transformer/encoder.py
index a6a85ae..987924f 100644
--- a/funasr/models/transformer/encoder.py
+++ b/funasr/models/transformer/encoder.py
@@ -64,7 +64,7 @@
         stochastic_depth_rate=0.0,
     ):
         """Construct an EncoderLayer object."""
-        super(EncoderLayer, self).__init__()
+        super().__init__()
         self.self_attn = self_attn
         self.feed_forward = feed_forward
         self.norm1 = LayerNorm(size)

--
Gitblit v1.9.1