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

Method forward_layers

OSUM/wenet/LLM/decoder.py:126–148  ·  view source on GitHub ↗
(
        self,
        xs: torch.Tensor,
        att_mask: torch.Tensor,
        pos_emb: torch.Tensor,
        kv_caches: Optional[List[T_CACHE]] = None,
    )

Source from the content-addressed store, hash-verified

124 return xs, kv_caches
125
126 def forward_layers(
127 self,
128 xs: torch.Tensor,
129 att_mask: torch.Tensor,
130 pos_emb: torch.Tensor,
131 kv_caches: Optional[List[T_CACHE]] = None,
132 ) -> Tuple[torch.Tensor, Union[List[T_CACHE], None]]:
133 if self.training:
134 for (i, layer) in enumerate(self.decoders):
135 xs, _, _, _ = layer(xs, att_mask, pos_emb)
136 new_kv_caches = kv_caches
137 else:
138 assert kv_caches is not None
139 new_kv_caches = []
140 for (i, layer) in enumerate(self.decoders):
141 xs, _, new_kv_cache, _ = layer(xs,
142 att_mask,
143 pos_emb,
144 att_cache=(kv_caches[i][0],
145 kv_caches[i][1]))
146 new_kv_caches.append(new_kv_cache)
147
148 return xs, new_kv_caches
149
150 @torch.jit.ignore(drop=True)
151 def forward_layers_checkpointed(self, xs: torch.Tensor,

Callers 1

forwardMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected