mirror of
https://github.com/wassname/peft.git
synced 2026-09-26 14:00:37 +08:00
33 lines
1.3 KiB
Python
33 lines
1.3 KiB
Python
import torch
|
|
|
|
|
|
# needed for prefix-tuning of bloom model
|
|
def bloom_model_postprocess_past_key_value(past_key_values):
|
|
past_key_values = torch.cat(past_key_values)
|
|
total_layers, batch_size, num_attention_heads, num_virtual_tokens, head_dim = past_key_values.shape
|
|
keys = past_key_values[: total_layers // 2]
|
|
keys = keys.transpose(2, 3).reshape(
|
|
total_layers // 2, batch_size * num_attention_heads, head_dim, num_virtual_tokens
|
|
)
|
|
values = past_key_values[total_layers // 2 :]
|
|
values = values.reshape(total_layers // 2, batch_size * num_attention_heads, num_virtual_tokens, head_dim)
|
|
|
|
return tuple(zip(keys, values))
|
|
|
|
|
|
# copied from transformers.models.bart.modeling_bart
|
|
def shift_tokens_right(input_ids: torch.Tensor, pad_token_id: int, decoder_start_token_id: int):
|
|
"""
|
|
Shift input ids one token to the right.
|
|
"""
|
|
shifted_input_ids = input_ids.new_zeros(input_ids.shape)
|
|
shifted_input_ids[:, 1:] = input_ids[:, :-1].clone()
|
|
shifted_input_ids[:, 0] = decoder_start_token_id
|
|
|
|
if pad_token_id is None:
|
|
raise ValueError("self.model.config.pad_token_id has to be defined.")
|
|
# replace possible -100 values in labels by `pad_token_id`
|
|
shifted_input_ids.masked_fill_(shifted_input_ids == -100, pad_token_id)
|
|
|
|
return shifted_input_ids
|