MCPcopy Create free account
hub / github.com/ASLP-lab/OSUM / __init__

Method __init__

OSUM/wenet/transformer/decoder.py:63–144  ·  view source on GitHub ↗
(
        self,
        vocab_size: int,
        encoder_output_size: int,
        attention_heads: int = 4,
        linear_units: int = 2048,
        num_blocks: int = 6,
        dropout_rate: float = 0.1,
        positional_dropout_rate: float = 0.1,
        self_attention_dropout_rate: float = 0.0,
        src_attention_dropout_rate: float = 0.0,
        input_layer: str = "embed",
        use_output_layer: bool = True,
        normalize_before: bool = True,
        src_attention: bool = True,
        query_bias: bool = True,
        key_bias: bool = True,
        value_bias: bool = True,
        activation_type: str = "relu",
        gradient_checkpointing: bool = False,
        tie_word_embedding: bool = False,
        use_sdpa: bool = False,
        layer_norm_type: str = 'layer_norm',
        norm_eps: float = 1e-5,
        n_kv_head: Optional[int] = None,
        head_dim: Optional[int] = None,
        mlp_type: str = 'position_wise_feed_forward',
        mlp_bias: bool = True,
        n_expert: int = 8,
        n_expert_activated: int = 2,
    )

Source from the content-addressed store, hash-verified

61 """
62
63 def __init__(
64 self,
65 vocab_size: int,
66 encoder_output_size: int,
67 attention_heads: int = 4,
68 linear_units: int = 2048,
69 num_blocks: int = 6,
70 dropout_rate: float = 0.1,
71 positional_dropout_rate: float = 0.1,
72 self_attention_dropout_rate: float = 0.0,
73 src_attention_dropout_rate: float = 0.0,
74 input_layer: str = "embed",
75 use_output_layer: bool = True,
76 normalize_before: bool = True,
77 src_attention: bool = True,
78 query_bias: bool = True,
79 key_bias: bool = True,
80 value_bias: bool = True,
81 activation_type: str = "relu",
82 gradient_checkpointing: bool = False,
83 tie_word_embedding: bool = False,
84 use_sdpa: bool = False,
85 layer_norm_type: str = 'layer_norm',
86 norm_eps: float = 1e-5,
87 n_kv_head: Optional[int] = None,
88 head_dim: Optional[int] = None,
89 mlp_type: str = 'position_wise_feed_forward',
90 mlp_bias: bool = True,
91 n_expert: int = 8,
92 n_expert_activated: int = 2,
93 ):
94 super().__init__()
95 attention_dim = encoder_output_size
96 activation = WENET_ACTIVATION_CLASSES[activation_type]()
97
98 self.embed = torch.nn.Sequential(
99 torch.nn.Identity() if input_layer == "no_pos" else
100 torch.nn.Embedding(vocab_size, attention_dim),
101 WENET_EMB_CLASSES[input_layer](attention_dim,
102 positional_dropout_rate),
103 )
104
105 assert layer_norm_type in ['layer_norm', 'rms_norm']
106 self.normalize_before = normalize_before
107 self.after_norm = WENET_NORM_CLASSES[layer_norm_type](attention_dim,
108 eps=norm_eps)
109 self.use_output_layer = use_output_layer
110 if use_output_layer:
111 self.output_layer = torch.nn.Linear(attention_dim, vocab_size)
112 else:
113 self.output_layer = torch.nn.Identity()
114 self.num_blocks = num_blocks
115
116 mlp_class = WENET_MLP_CLASSES[mlp_type]
117 self.decoders = torch.nn.ModuleList([
118 DecoderLayer(
119 attention_dim,
120 WENET_ATTENTION_CLASSES["selfattn"](

Callers 1

__init__Method · 0.45

Calls 1

DecoderLayerClass · 0.90

Tested by

no test coverage detected