2023-09-09 12:00:29 +08:00
|
|
|
import torch
|
2023-09-14 08:34:30 +08:00
|
|
|
import intel_extension_for_pytorch as ipex # pylint: disable=import-error, unused-import
|
2023-09-09 12:00:29 +08:00
|
|
|
|
|
|
|
# pylint: disable=protected-access, missing-function-docstring, line-too-long
|
|
|
|
|
|
|
|
original_torch_bmm = torch.bmm
|
2023-09-14 08:34:30 +08:00
|
|
|
|
|
|
|
|
2023-09-09 12:00:29 +08:00
|
|
|
def torch_bmm(input, mat2, *, out=None):
|
|
|
|
if input.dtype != mat2.dtype:
|
|
|
|
mat2 = mat2.to(input.dtype)
|
|
|
|
|
2023-09-14 08:34:30 +08:00
|
|
|
# ARC GPUs can't allocate more than 4GB to a single block, Slice it:
|
|
|
|
batch_size_attention, input_tokens, mat2_shape = (
|
|
|
|
input.shape[0],
|
|
|
|
input.shape[1],
|
|
|
|
mat2.shape[2],
|
|
|
|
)
|
2023-10-06 17:14:33 +08:00
|
|
|
block_multiply = input.element_size()
|
|
|
|
slice_block_size = input_tokens * mat2_shape / 1024 / 1024 * block_multiply
|
|
|
|
block_size = batch_size_attention * slice_block_size
|
|
|
|
|
2023-09-09 12:00:29 +08:00
|
|
|
split_slice_size = batch_size_attention
|
2023-10-06 17:14:33 +08:00
|
|
|
if block_size > 4:
|
2023-09-09 12:00:29 +08:00
|
|
|
do_split = True
|
2023-09-14 08:34:30 +08:00
|
|
|
# Find something divisible with the input_tokens
|
2023-10-06 17:14:33 +08:00
|
|
|
while (split_slice_size * slice_block_size) > 4:
|
2023-09-09 12:00:29 +08:00
|
|
|
split_slice_size = split_slice_size // 2
|
|
|
|
if split_slice_size <= 1:
|
|
|
|
split_slice_size = 1
|
|
|
|
break
|
|
|
|
else:
|
|
|
|
do_split = False
|
|
|
|
|
|
|
|
split_2_slice_size = input_tokens
|
2023-10-06 17:14:33 +08:00
|
|
|
if split_slice_size * slice_block_size > 4:
|
|
|
|
slice_block_size2 = split_slice_size * mat2_shape / 1024 / 1024 * block_multiply
|
2023-09-09 12:00:29 +08:00
|
|
|
do_split_2 = True
|
2023-09-14 08:34:30 +08:00
|
|
|
# Find something divisible with the input_tokens
|
2023-10-06 17:14:33 +08:00
|
|
|
while (split_2_slice_size * slice_block_size2) > 4:
|
2023-09-09 12:00:29 +08:00
|
|
|
split_2_slice_size = split_2_slice_size // 2
|
|
|
|
if split_2_slice_size <= 1:
|
|
|
|
split_2_slice_size = 1
|
|
|
|
break
|
|
|
|
else:
|
|
|
|
do_split_2 = False
|
|
|
|
|
|
|
|
if do_split:
|
2023-09-14 08:34:30 +08:00
|
|
|
hidden_states = torch.zeros(
|
|
|
|
input.shape[0],
|
|
|
|
input.shape[1],
|
|
|
|
mat2.shape[2],
|
|
|
|
device=input.device,
|
|
|
|
dtype=input.dtype,
|
|
|
|
)
|
2023-09-09 12:00:29 +08:00
|
|
|
for i in range(batch_size_attention // split_slice_size):
|
|
|
|
start_idx = i * split_slice_size
|
|
|
|
end_idx = (i + 1) * split_slice_size
|
|
|
|
if do_split_2:
|
2023-09-14 08:34:30 +08:00
|
|
|
for i2 in range(
|
|
|
|
input_tokens // split_2_slice_size
|
|
|
|
): # pylint: disable=invalid-name
|
2023-09-09 12:00:29 +08:00
|
|
|
start_idx_2 = i2 * split_2_slice_size
|
|
|
|
end_idx_2 = (i2 + 1) * split_2_slice_size
|
2023-09-14 08:34:30 +08:00
|
|
|
hidden_states[
|
|
|
|
start_idx:end_idx, start_idx_2:end_idx_2
|
|
|
|
] = original_torch_bmm(
|
2023-09-09 12:00:29 +08:00
|
|
|
input[start_idx:end_idx, start_idx_2:end_idx_2],
|
|
|
|
mat2[start_idx:end_idx, start_idx_2:end_idx_2],
|
2023-09-14 08:34:30 +08:00
|
|
|
out=out,
|
2023-09-09 12:00:29 +08:00
|
|
|
)
|
|
|
|
else:
|
|
|
|
hidden_states[start_idx:end_idx] = original_torch_bmm(
|
2023-09-14 08:34:30 +08:00
|
|
|
input[start_idx:end_idx], mat2[start_idx:end_idx], out=out
|
2023-09-09 12:00:29 +08:00
|
|
|
)
|
|
|
|
else:
|
|
|
|
return original_torch_bmm(input, mat2, out=out)
|
|
|
|
return hidden_states
|
|
|
|
|
2023-09-14 08:34:30 +08:00
|
|
|
|
2023-09-09 12:00:29 +08:00
|
|
|
original_scaled_dot_product_attention = torch.nn.functional.scaled_dot_product_attention
|
2023-09-14 08:34:30 +08:00
|
|
|
|
|
|
|
|
|
|
|
def scaled_dot_product_attention(
|
|
|
|
query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False
|
|
|
|
):
|
|
|
|
# ARC GPUs can't allocate more than 4GB to a single block, Slice it:
|
2023-10-06 17:14:33 +08:00
|
|
|
if len(query.shape) == 3:
|
|
|
|
batch_size_attention, query_tokens, shape_four = query.shape
|
|
|
|
shape_one = 1
|
|
|
|
no_shape_one = True
|
|
|
|
else:
|
|
|
|
shape_one, batch_size_attention, query_tokens, shape_four = query.shape
|
|
|
|
no_shape_one = False
|
|
|
|
|
|
|
|
block_multiply = query.element_size()
|
|
|
|
slice_block_size = (
|
|
|
|
shape_one * query_tokens * shape_four / 1024 / 1024 * block_multiply
|
|
|
|
)
|
|
|
|
block_size = batch_size_attention * slice_block_size
|
|
|
|
|
2023-09-09 12:00:29 +08:00
|
|
|
split_slice_size = batch_size_attention
|
2023-10-06 17:14:33 +08:00
|
|
|
if block_size > 4:
|
2023-09-09 12:00:29 +08:00
|
|
|
do_split = True
|
2023-09-14 08:34:30 +08:00
|
|
|
# Find something divisible with the shape_one
|
2023-10-06 17:14:33 +08:00
|
|
|
while (split_slice_size * slice_block_size) > 4:
|
2023-09-09 12:00:29 +08:00
|
|
|
split_slice_size = split_slice_size // 2
|
|
|
|
if split_slice_size <= 1:
|
|
|
|
split_slice_size = 1
|
|
|
|
break
|
|
|
|
else:
|
|
|
|
do_split = False
|
|
|
|
|
|
|
|
split_2_slice_size = query_tokens
|
2023-10-06 17:14:33 +08:00
|
|
|
if split_slice_size * slice_block_size > 4:
|
|
|
|
slice_block_size2 = (
|
|
|
|
shape_one * split_slice_size * shape_four / 1024 / 1024 * block_multiply
|
|
|
|
)
|
2023-09-09 12:00:29 +08:00
|
|
|
do_split_2 = True
|
2023-09-14 08:34:30 +08:00
|
|
|
# Find something divisible with the batch_size_attention
|
2023-10-06 17:14:33 +08:00
|
|
|
while (split_2_slice_size * slice_block_size2) > 4:
|
2023-09-09 12:00:29 +08:00
|
|
|
split_2_slice_size = split_2_slice_size // 2
|
|
|
|
if split_2_slice_size <= 1:
|
|
|
|
split_2_slice_size = 1
|
|
|
|
break
|
|
|
|
else:
|
|
|
|
do_split_2 = False
|
|
|
|
|
|
|
|
if do_split:
|
|
|
|
hidden_states = torch.zeros(query.shape, device=query.device, dtype=query.dtype)
|
|
|
|
for i in range(batch_size_attention // split_slice_size):
|
|
|
|
start_idx = i * split_slice_size
|
|
|
|
end_idx = (i + 1) * split_slice_size
|
|
|
|
if do_split_2:
|
2023-09-14 08:34:30 +08:00
|
|
|
for i2 in range(
|
|
|
|
query_tokens // split_2_slice_size
|
|
|
|
): # pylint: disable=invalid-name
|
2023-09-09 12:00:29 +08:00
|
|
|
start_idx_2 = i2 * split_2_slice_size
|
|
|
|
end_idx_2 = (i2 + 1) * split_2_slice_size
|
2023-10-06 17:14:33 +08:00
|
|
|
if no_shape_one:
|
|
|
|
hidden_states[
|
|
|
|
start_idx:end_idx, start_idx_2:end_idx_2
|
|
|
|
] = original_scaled_dot_product_attention(
|
|
|
|
query[start_idx:end_idx, start_idx_2:end_idx_2],
|
|
|
|
key[start_idx:end_idx, start_idx_2:end_idx_2],
|
|
|
|
value[start_idx:end_idx, start_idx_2:end_idx_2],
|
|
|
|
attn_mask=attn_mask[
|
|
|
|
start_idx:end_idx, start_idx_2:end_idx_2
|
|
|
|
]
|
|
|
|
if attn_mask is not None
|
|
|
|
else attn_mask,
|
|
|
|
dropout_p=dropout_p,
|
|
|
|
is_causal=is_causal,
|
|
|
|
)
|
|
|
|
else:
|
|
|
|
hidden_states[
|
|
|
|
:, start_idx:end_idx, start_idx_2:end_idx_2
|
|
|
|
] = original_scaled_dot_product_attention(
|
|
|
|
query[:, start_idx:end_idx, start_idx_2:end_idx_2],
|
|
|
|
key[:, start_idx:end_idx, start_idx_2:end_idx_2],
|
|
|
|
value[:, start_idx:end_idx, start_idx_2:end_idx_2],
|
|
|
|
attn_mask=attn_mask[
|
|
|
|
:, start_idx:end_idx, start_idx_2:end_idx_2
|
|
|
|
]
|
|
|
|
if attn_mask is not None
|
|
|
|
else attn_mask,
|
|
|
|
dropout_p=dropout_p,
|
|
|
|
is_causal=is_causal,
|
|
|
|
)
|
|
|
|
else:
|
|
|
|
if no_shape_one:
|
2023-09-14 08:34:30 +08:00
|
|
|
hidden_states[
|
2023-10-06 17:14:33 +08:00
|
|
|
start_idx:end_idx
|
2023-09-14 08:34:30 +08:00
|
|
|
] = original_scaled_dot_product_attention(
|
2023-10-06 17:14:33 +08:00
|
|
|
query[start_idx:end_idx],
|
|
|
|
key[start_idx:end_idx],
|
|
|
|
value[start_idx:end_idx],
|
|
|
|
attn_mask=attn_mask[start_idx:end_idx]
|
|
|
|
if attn_mask is not None
|
|
|
|
else attn_mask,
|
|
|
|
dropout_p=dropout_p,
|
|
|
|
is_causal=is_causal,
|
|
|
|
)
|
|
|
|
else:
|
|
|
|
hidden_states[
|
|
|
|
:, start_idx:end_idx
|
|
|
|
] = original_scaled_dot_product_attention(
|
|
|
|
query[:, start_idx:end_idx],
|
|
|
|
key[:, start_idx:end_idx],
|
|
|
|
value[:, start_idx:end_idx],
|
|
|
|
attn_mask=attn_mask[:, start_idx:end_idx]
|
2023-09-14 08:34:30 +08:00
|
|
|
if attn_mask is not None
|
|
|
|
else attn_mask,
|
|
|
|
dropout_p=dropout_p,
|
|
|
|
is_causal=is_causal,
|
2023-09-09 12:00:29 +08:00
|
|
|
)
|
|
|
|
else:
|
|
|
|
return original_scaled_dot_product_attention(
|
2023-09-14 08:34:30 +08:00
|
|
|
query,
|
|
|
|
key,
|
|
|
|
value,
|
|
|
|
attn_mask=attn_mask,
|
|
|
|
dropout_p=dropout_p,
|
|
|
|
is_causal=is_causal,
|
2023-09-09 12:00:29 +08:00
|
|
|
)
|
|
|
|
return hidden_states
|
|
|
|
|
2023-09-14 08:34:30 +08:00
|
|
|
|
2023-09-09 12:00:29 +08:00
|
|
|
def attention_init():
|
2023-09-14 08:34:30 +08:00
|
|
|
# ARC GPUs can't allocate more than 4GB to a single block:
|
2023-09-09 12:00:29 +08:00
|
|
|
torch.bmm = torch_bmm
|
|
|
|
torch.nn.functional.scaled_dot_product_attention = scaled_dot_product_attention
|