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

Method forward

OSUM/wenet/LLM/decoder.py:106–124  ·  view source on GitHub ↗
(
        self,
        input: torch.Tensor,
        att_mask: torch.Tensor,
        input_position: Union[int, torch.Tensor] = 0,
        kv_caches: Optional[List[T_CACHE]] = None,
    )

Source from the content-addressed store, hash-verified

104 self.gradient_checkpointing = gradient_checkpointing
105
106 def forward(
107 self,
108 input: torch.Tensor,
109 att_mask: torch.Tensor,
110 input_position: Union[int, torch.Tensor] = 0,
111 kv_caches: Optional[List[T_CACHE]] = None,
112 ) -> Tuple[torch.Tensor, Union[List[T_CACHE], None]]:
113 xs, pos_emb = self.pos_enc(input, offset=input_position)
114 if self.use_sdpa:
115 att_mask = mask_to_bias(att_mask, xs.dtype)
116
117 if self.gradient_checkpointing and self.training:
118 xs = self.forward_layers_checkpointed(xs, att_mask, pos_emb)
119 else:
120 xs, kv_caches = self.forward_layers(xs, att_mask, pos_emb,
121 kv_caches)
122 if self.pre_norm and self.final_norm is not None:
123 xs = self.final_norm(xs)
124 return xs, kv_caches
125
126 def forward_layers(
127 self,

Callers

nothing calls this directly

Calls 3

forward_layersMethod · 0.95
mask_to_biasFunction · 0.90

Tested by

no test coverage detected