MCPcopy Create free account
hub / github.com/OpenMOSS/MOSS / forward

Method forward

models/modeling_moss.py:151–226  ·  view source on GitHub ↗
(
        self,
        hidden_states: Optional[torch.FloatTensor],
        layer_past: Optional[Tuple[torch.Tensor]] = None,
        attention_mask: Optional[torch.FloatTensor] = None,
        position_ids: Optional[torch.LongTensor] = None,
        head_mask: Optional[torch.FloatTensor] = None,
        use_cache: Optional[bool] = False,
        output_attentions: Optional[bool] = False,
    )

Source from the content-addressed store, hash-verified

149 return attn_output, attn_weights
150
151 def forward(
152 self,
153 hidden_states: Optional[torch.FloatTensor],
154 layer_past: Optional[Tuple[torch.Tensor]] = None,
155 attention_mask: Optional[torch.FloatTensor] = None,
156 position_ids: Optional[torch.LongTensor] = None,
157 head_mask: Optional[torch.FloatTensor] = None,
158 use_cache: Optional[bool] = False,
159 output_attentions: Optional[bool] = False,
160 ) -> Union[
161 Tuple[torch.Tensor, Tuple[torch.Tensor]],
162 Optional[Tuple[torch.Tensor, Tuple[torch.Tensor], Tuple[torch.Tensor, ...]]],
163 ]:
164 qkv = self.qkv_proj(hidden_states)
165 # TODO(enijkamp): factor out number of logical TPU-v4 cores or make forward pass agnostic
166 mp_num = 4
167 qkv_split = qkv.reshape(qkv.shape[:-1] + (mp_num, -1))
168
169 local_dim = self.head_dim * self.num_attention_heads // mp_num
170 query, value, key = torch.split(qkv_split, local_dim, dim=-1)
171 query = self._split_heads(query, self.num_attention_heads, self.head_dim, mp_num=mp_num)
172 key = self._split_heads(key, self.num_attention_heads, self.head_dim, mp_num=mp_num)
173
174 value = self._split_heads(value, self.num_attention_heads, self.head_dim, mp_num=mp_num)
175 value = value.permute(0, 2, 1, 3)
176
177 embed_positions = self.embed_positions
178 if embed_positions.device != position_ids.device:
179 embed_positions = embed_positions.to(position_ids.device)
180 self.embed_positions = embed_positions
181
182 sincos = embed_positions[position_ids]
183 sin, cos = torch.split(sincos, sincos.shape[-1] // 2, dim=-1)
184
185 if self.rotary_dim is not None:
186 k_rot = key[:, :, :, : self.rotary_dim]
187 k_pass = key[:, :, :, self.rotary_dim :]
188
189 q_rot = query[:, :, :, : self.rotary_dim]
190 q_pass = query[:, :, :, self.rotary_dim :]
191
192 k_rot = apply_rotary_pos_emb(k_rot, sin, cos)
193 q_rot = apply_rotary_pos_emb(q_rot, sin, cos)
194
195 key = torch.cat([k_rot, k_pass], dim=-1)
196 query = torch.cat([q_rot, q_pass], dim=-1)
197 else:
198 key = apply_rotary_pos_emb(key, sin, cos)
199 query = apply_rotary_pos_emb(query, sin, cos)
200
201 key = key.permute(0, 2, 1, 3)
202 query = query.permute(0, 2, 1, 3)
203
204 if layer_past is not None:
205 past_key = layer_past[0]
206 past_value = layer_past[1]
207 key = torch.cat((past_key, key), dim=-2)
208 value = torch.cat((past_value, value), dim=-2)

Callers

nothing calls this directly

Calls 4

_split_headsMethod · 0.95
_attnMethod · 0.95
_merge_headsMethod · 0.95
apply_rotary_pos_embFunction · 0.70

Tested by

no test coverage detected