shixian.shi
2024-01-10 2d0c8274a6dd0db9d5f7a00e71401bf7b4d65553
funasr/models/data2vec/data2vec_encoder.py
@@ -11,7 +11,6 @@
import torch.nn as nn
import torch.nn.functional as F
from funasr.models.encoder.abs_encoder import AbsEncoder
from funasr.models.data2vec.data_utils import compute_mask_indices
from funasr.models.data2vec.ema_module import EMAModule
from funasr.models.data2vec.grad_multiply import GradMultiply
@@ -28,7 +27,7 @@
    return end - r * pct_remaining
class Data2VecEncoder(AbsEncoder):
class Data2VecEncoder(nn.Module):
    def __init__(
            self,
            # for ConvFeatureExtractionModel