From 9b4e9cc8a0311e5243d69b73ed073e7ea441982e Mon Sep 17 00:00:00 2001
From: 游雁 <zhifu.gzf@alibaba-inc.com>
Date: 星期三, 27 三月 2024 16:05:29 +0800
Subject: [PATCH] train update
---
funasr/models/rwkv_bat/rwkv_encoder.py | 30 +++++++++++++++++-------------
1 files changed, 17 insertions(+), 13 deletions(-)
diff --git a/funasr/models/rwkv_bat/rwkv_encoder.py b/funasr/models/rwkv_bat/rwkv_encoder.py
index 3454299..c0e5f42 100644
--- a/funasr/models/rwkv_bat/rwkv_encoder.py
+++ b/funasr/models/rwkv_bat/rwkv_encoder.py
@@ -1,17 +1,20 @@
-"""RWKV encoder definition for Transducer models."""
-
-import math
-from typing import Dict, List, Optional, Tuple
+#!/usr/bin/env python3
+# -*- encoding: utf-8 -*-
+# Copyright FunASR (https://github.com/alibaba-damo-academy/FunASR). All Rights Reserved.
+# MIT License (https://opensource.org/licenses/MIT)
import torch
+from typing import Dict, List, Optional, Tuple
-from funasr.models.encoder.abs_encoder import AbsEncoder
+from funasr.register import tables
from funasr.models.rwkv_bat.rwkv import RWKV
from funasr.models.transformer.layer_norm import LayerNorm
-from funasr.models.rwkv_bat.rwkv_subsampling import RWKVConvInput
from funasr.models.transformer.utils.nets_utils import make_source_mask
+from funasr.models.rwkv_bat.rwkv_subsampling import RWKVConvInput
-class RWKVEncoder(AbsEncoder):
+
+@tables.register("encoder_classes", "RWKVEncoder")
+class RWKVEncoder(torch.nn.Module):
"""RWKV encoder module.
Based on https://arxiv.org/pdf/2305.13048.pdf.
@@ -44,6 +47,7 @@
subsampling_factor: int =4,
time_reduction_factor: int = 1,
kernel: int = 3,
+ **kwargs,
) -> None:
"""Construct a RWKVEncoder object."""
super().__init__()
@@ -113,12 +117,12 @@
x = self.embed_norm(x)
olens = mask.eq(0).sum(1)
- # for training
- # for block in self.rwkv_blocks:
- # x, _ = block(x)
-
- # for streaming inference
- x = self.rwkv_infer(x)
+ if self.training:
+ for block in self.rwkv_blocks:
+ x, _ = block(x)
+ else:
+ x = self.rwkv_infer(x)
+
x = self.final_norm(x)
if self.time_reduction_factor > 1:
--
Gitblit v1.9.1