MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / forward

Method forward

SwissArmyTransformer/sat/model/transformer.py:571–749  ·  view source on GitHub ↗
(self, input_ids, position_ids, attention_mask, *,
                output_hidden_states=False, **kw_args)

Source from the content-addressed store, hash-verified

569 self.final_layernorm = layernorm(hidden_size, eps=layernorm_epsilon)
570
571 def forward(self, input_ids, position_ids, attention_mask, *,
572 output_hidden_states=False, **kw_args):
573 # sanity check
574 assert len(input_ids.shape) >= 2
575 batch_size, query_length = input_ids.shape[:2]
576
577 if attention_mask is None:
578 # Definition: None means full attention
579 attention_mask = torch.ones(1, 1, device=input_ids.device)
580 elif isinstance(attention_mask, int) and (attention_mask < 0):
581 # Definition: -1 means lower triangular attention mask
582 attention_mask = torch.ones(query_length, query_length,
583 device=input_ids.device).tril()
584
585 attention_mask = attention_mask.type_as(
586 next(self.parameters())
587 )
588 assert len(attention_mask.shape) == 2 or \
589 len(attention_mask.shape) == 4 and attention_mask.shape[1] == 1
590
591 # initial output_cross_layer might be generated by word/position_embedding_forward
592 output_cross_layer = {}
593
594 # embedding part
595 if 'word_embedding_forward' in self.hooks:
596 hidden_states = self.hooks['word_embedding_forward'](input_ids, output_cross_layer=output_cross_layer, **kw_args)
597 else: # default
598 hidden_states = HOOKS_DEFAULT['word_embedding_forward'](self, input_ids, output_cross_layer=output_cross_layer,**kw_args)
599
600 # handle position embedding
601 if 'position_embedding_forward' in self.hooks:
602 position_embeddings = self.hooks['position_embedding_forward'](position_ids, output_cross_layer=output_cross_layer, **kw_args)
603 else:
604 assert len(position_ids.shape) <= 2
605 assert position_ids.shape[-1] == hidden_states.shape[1], (position_ids.shape, hidden_states.shape)
606 position_embeddings = HOOKS_DEFAULT['position_embedding_forward'](self, position_ids, output_cross_layer=output_cross_layer, **kw_args)
607 if position_embeddings is not None:
608 # import pdb; pdb.set_trace() # !! DEBUG
609 # hs: torch.Size([2, 88034, 1920])
610 # pos: torch.Size([1, 23522, 1920])
611 hidden_states = hidden_states + position_embeddings
612
613 hidden_states = self.embedding_dropout(hidden_states)
614
615 output_per_layers = []
616 if self.checkpoint_activations:
617 # define custom_forward for checkpointing
618 def custom(start, end, kw_args_index, cross_layer_index):
619 def custom_forward(*inputs):
620 layers_ = self.layers[start:end]
621 x_, mask = inputs[0], inputs[1]
622
623 # recover kw_args and output_cross_layer
624 flat_inputs = inputs[2:]
625 kw_args, output_cross_layer = {}, {}
626 for k, idx in kw_args_index.items():
627 kw_args[k] = flat_inputs[idx]
628 for k, idx in cross_layer_index.items():

Callers

nothing calls this directly

Calls 5

parametersMethod · 0.80
appendMethod · 0.80
extendMethod · 0.80
checkpointFunction · 0.50

Tested by

no test coverage detected