MCPcopy Create free account
hub / github.com/THUDM/LongWriter / get_masks

Method get_masks

train/patch/modeling_chatglm.py:588–604  ·  view source on GitHub ↗
(self, input_ids, past_key_values, padding_mask=None)

Source from the content-addressed store, hash-verified

586 return
587
588 def get_masks(self, input_ids, past_key_values, padding_mask=None):
589 batch_size, seq_length = input_ids.shape
590 full_attention_mask = torch.ones(batch_size, seq_length, seq_length, device=input_ids.device)
591 full_attention_mask.tril_()
592 past_length = 0
593 if past_key_values:
594 past_length = past_key_values[0][0].shape[0]
595 if past_length:
596 full_attention_mask = torch.cat((torch.ones(batch_size, seq_length, past_length,
597 device=input_ids.device), full_attention_mask), dim=-1)
598 if padding_mask is not None:
599 full_attention_mask = full_attention_mask * padding_mask.unsqueeze(1)
600 if not past_length and padding_mask is not None:
601 full_attention_mask -= padding_mask.unsqueeze(-1) - 1
602 full_attention_mask = (full_attention_mask < 0.5).bool()
603 full_attention_mask.unsqueeze_(1)
604 return full_attention_mask
605
606 def get_position_ids(self, input_ids, device):
607 batch_size, seq_length = input_ids.shape

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected