From f7ef4b5e3133bde6dc4d3ae7448f90e3b3c0892e Mon Sep 17 00:00:00 2001 From: Reza Yazdani Date: Thu, 28 Oct 2021 02:06:42 +0500 Subject: [PATCH 1/3] fixing the softmax masking when using triangular masking --- .../transformer/inference/csrc/pt_binding.cpp | 13 ++-- csrc/transformer/inference/csrc/softmax.cu | 43 +++++++------ .../inference/transformer_inference.py | 61 +++++++++++++------ 3 files changed, 69 insertions(+), 48 deletions(-) diff --git a/csrc/transformer/inference/csrc/pt_binding.cpp b/csrc/transformer/inference/csrc/pt_binding.cpp index 342fbe7b253c..1ebadaeb53b4 100644 --- a/csrc/transformer/inference/csrc/pt_binding.cpp +++ b/csrc/transformer/inference/csrc/pt_binding.cpp @@ -11,7 +11,7 @@ std::array gemm_algos = std::array({99, 99, 99}); template at::Tensor ds_softmax(at::Tensor& attn_scores, - T* attn_mask_ptr, + at::Tensor& attn_mask, bool triangular, bool recompute, bool local_attention, @@ -22,9 +22,8 @@ at::Tensor ds_softmax(at::Tensor& attn_scores, int seq_len = attn_scores_c.size(2); int soft_len = attn_scores_c.size(3); int heads = attn_scores_c.size(1); - launch_attn_softmax_v2((T*)attn_scores_c.data_ptr(), - attn_mask_ptr, + (attn_mask.sizes().size() > 1 ? (T*)attn_mask.data_ptr() : nullptr), triangular, recompute, local_attention, @@ -42,7 +41,7 @@ at::Tensor ds_softmax(at::Tensor& attn_scores, template void attention_unfused(at::Tensor& prev_key_cont, at::Tensor& query_cont, - T* attn_mask_ptr, + at::Tensor& attn_mask, at::Tensor& prev_value_cont, at::Tensor& output, int& bsz, @@ -81,8 +80,8 @@ void attention_unfused(at::Tensor& prev_key_cont, seq_len * soft_len, bsz * heads, CUBLAS_GEMM_DEFAULT_TENSOR_OP); - attn_score = ds_softmax( - attn_score, attn_mask_ptr, triangular, recompute, local_attention, window_size); + attn_score = + ds_softmax(attn_score, attn_mask, triangular, recompute, local_attention, window_size); alpha = 1.0; cublas_strided_batched_gemm(Context::Instance().GetCublasHandle(), k, @@ -139,7 +138,7 @@ std::vector ds_softmax_context(at::Tensor& query, at::empty({prev_value.size(0), heads, seq_len, prev_value.size(2) / heads}, options); attention_unfused(prev_key_cont, query_cont, - (no_masking ? nullptr : (T*)attn_mask.data_ptr()), + attn_mask, //(no_masking ? nullptr : (T*)attn_mask.data_ptr()), prev_value_cont, output, bsz, diff --git a/csrc/transformer/inference/csrc/softmax.cu b/csrc/transformer/inference/csrc/softmax.cu index 3ffad01b623a..950ae6aeaafb 100644 --- a/csrc/transformer/inference/csrc/softmax.cu +++ b/csrc/transformer/inference/csrc/softmax.cu @@ -9,7 +9,7 @@ #define ATTN_THREADS 1024 #define MAX_REG_SIZE 8 -#define minus_infinity (-1 * std::numeric_limits::infinity()) +#define minus_infinity -10000.0 void CheckCudaErrorAux(const char* file, unsigned line) { @@ -94,10 +94,10 @@ __global__ void attn_softmax_v2(__half* vals, (data_id + 3) > window_stride) ? __half2float(vals[data_id + 3]) : minus_infinity; - if (mask && !triangular && recompute) { + if (mask && recompute) { low_data[i].x += __half2float(mask[data_id + mask_offset]); low_data[i].y += __half2float(mask[data_id + mask_offset + 1]); - high_data[i].y += __half2float(mask[data_id + mask_offset + 2]); + high_data[i].x += __half2float(mask[data_id + mask_offset + 2]); high_data[i].y += __half2float(mask[data_id + mask_offset + 3]); } } else { @@ -114,15 +114,15 @@ __global__ void attn_softmax_v2(__half* vals, ? __half2float(vals[data_id + 2]) : minus_infinity; high_data[i].y = minus_infinity; - if (mask && !triangular && recompute) { + if (mask && recompute) { low_data[i].x += __half2float(mask[data_id + mask_offset]); if ((data_id + 1) < sequence_length) low_data[i].y += __half2float(mask[data_id + mask_offset + 1]); if ((data_id + 2) < sequence_length) high_data[i].x += __half2float(mask[data_id + mask_offset + 2]); - // high_data[i].y += __half2float(mask[data_id + mask_offset + 3]); } } + // if(lane == 0) printf("%f , %d, %d \n", low_data[i].x, data_id, seq_id); max_val = (low_data[i].x > max_val ? low_data[i].x : max_val); max_val = (low_data[i].y > max_val ? low_data[i].y : max_val); max_val = (high_data[i].x > max_val ? high_data[i].x : max_val); @@ -155,7 +155,6 @@ __global__ void attn_softmax_v2(__half* vals, max_val = g.shfl(max_val, threadIdx.x / WARP_SIZE); } - float sum = 0; for (int i = 0; i < iterations; i++) { low_data[i].x = __expf(low_data[i].x - max_val); @@ -181,7 +180,6 @@ __global__ void attn_softmax_v2(__half* vals, sum = g.shfl(sum, threadIdx.x / WARP_SIZE); } sum += 1e-6; - for (int i = 0; i < iterations; i++) { int data_id = i * (reduceWidth << 2) + (seq_lane << 2); @@ -265,7 +263,7 @@ __global__ void attn_softmax_v2(float* vals, (data_id + 3) > window_stride) ? vals[data_id + 3] : minus_infinity; - if (attn_mask && !triangular && recompute) { + if (attn_mask && recompute) { data[i].x += attn_mask[data_id + mask_offset]; data[i].y += attn_mask[data_id + mask_offset + 1]; data[i].z += attn_mask[data_id + mask_offset + 2]; @@ -282,7 +280,7 @@ __global__ void attn_softmax_v2(float* vals, ? (vals[data_id + 2]) : minus_infinity; data[i].w = minus_infinity; - if (attn_mask && !triangular && recompute) { + if (attn_mask && recompute) { data[i].x += attn_mask[data_id + mask_offset]; if ((data_id + 1) < sequence_length) data[i].y += attn_mask[data_id + mask_offset + 1]; @@ -390,19 +388,20 @@ void launch_attn_softmax_v2(T* vals, const int iterations = (sequence_length - 1) / (reduce_width << 2) + 1; if (sequence_length <= 32768) - attn_softmax_v2<<>>(vals, - mask, - triangular, - recompute, - local_attention, - window_size, - total_count, - heads, - sequence_length, - num_seq, - scale, - iterations, - reduce_width); + attn_softmax_v2<<>>( + vals, + mask, + triangular, + recompute, + local_attention, + window_size, + total_count, + (triangular ? (heads * batch_size) : heads), + sequence_length, + num_seq, + scale, + iterations, + reduce_width); else throw std::runtime_error("Unsupport Seq_Length!"); } diff --git a/deepspeed/ops/transformer/inference/transformer_inference.py b/deepspeed/ops/transformer/inference/transformer_inference.py index 218406c86b42..c398c3732924 100755 --- a/deepspeed/ops/transformer/inference/transformer_inference.py +++ b/deepspeed/ops/transformer/inference/transformer_inference.py @@ -159,11 +159,36 @@ def _transpose_for_context(x): (hidden_size_per_partition,) return x.view(*new_x_layer_shape) + def backup_attention(mixed_query, key_layer, value_layer, input_mask): + if layer_past is not None: + past_key, past_value = layer_past + key_layer = torch.cat((past_key.type_as(key_layer), key_layer), dim=-2) + value_layer = torch.cat((past_value.type_as(value_layer), + value_layer), + dim=-2) + query = _transpose_for_scores(mixed_query, False) + key = _transpose_for_scores(key_layer, True) + value = _transpose_for_scores(value_layer, False) + p = torch.matmul(query, key) + + ds_softmax = inference_cuda_module.softmax_fp16 if config.fp16 else \ + inference_cuda_module.softmax_fp32 + p = ds_softmax(p / (float(key.size(-2))**0.5), + input_mask, + True, + False, + False, + 256) + p = p.to(value.dtype) + context_layer = torch.matmul(p, value) + context_layer = _transpose_for_context(context_layer) + return context_layer, key_layer, value_layer + def compute_attention(qkv_out, input_mask): - score_context_func = inference_cuda_module.softmax_context_fp32 if (not config.fp16 or not config.triangular_masking) else \ + score_context_func = inference_cuda_module.softmax_context_fp32 if (not config.fp16) else \ inference_cuda_module.softmax_context_fp16 - if not config.triangular_masking: - qkv_out = qkv_out.float() + #if not config.triangular_masking: + # qkv_out = qkv_out.float() if merge_count > 0 and config.q_int8: split_dim = (qkv_out.dim() - 1) @@ -187,9 +212,14 @@ def compute_attention(qkv_out, input_mask): value_layer) = torch.split(qkv_out, (qkv_out.shape[-1] // 3), dim=(qkv_out.dim() - 1)) - + no_masking = input_mask is None + if no_masking: + input_mask = torch.empty(1) head_size = (mixed_query.shape[-1] // num_attention_heads_per_partition) - + #return backup_attention(mixed_query, + # key_layer, + # value_layer, + # input_mask) unfused_mode = not config.specialized_mode or \ mixed_query.shape[1] >= 32 or head_size > 128 @@ -210,17 +240,12 @@ def compute_attention(qkv_out, input_mask): True) / (norm_factor if config.scale_attention else 1.0) value_layer1 = _transpose_for_scores(value_layer, False, True) - no_masking = input_mask is None - if no_masking: - input_mask = torch.empty(1) - if layer_past is None: attn_key_value = score_context_func( mixed_query, (key_layer1 if unfused_mode else key_layer), torch.empty(1), - (input_mask - if config.triangular_masking or no_masking else input_mask.float()), + (input_mask), (value_layer1 if unfused_mode else value_layer), torch.empty(1), num_attention_heads_per_partition, @@ -235,8 +260,7 @@ def compute_attention(qkv_out, input_mask): mixed_query, (key_layer1 if unfused_mode else past_key.type_as(key_layer)), (key_layer1 if unfused_mode else key_layer), - (input_mask - if config.triangular_masking or no_masking else input_mask.float()), + (input_mask), (value_layer1 if unfused_mode else past_value.type_as(value_layer)), (value_layer1 if unfused_mode else value_layer), num_attention_heads_per_partition, @@ -254,8 +278,8 @@ def compute_attention(qkv_out, input_mask): # Transpose Context context_layer = _transpose_for_context(context_layer) - if (config.fp16 or config.q_int8) and not config.triangular_masking: - context_layer = context_layer.half() + #if (config.fp16 or config.q_int8) and not config.triangular_masking: + # context_layer = context_layer.half() return context_layer, key_layer, value_layer @@ -552,14 +576,13 @@ def __init__(self, self.config = config self.config.layer_id = DeepSpeedTransformerInference.layer_id DeepSpeedTransformerInference.layer_id += 1 - - self.attention = DeepSpeedSelfAttention(config, + self.attention = DeepSpeedSelfAttention(self.config, mp_group, quantize_scales, quantize_groups, merge_count, qkv_merging) - self.mlp = DeepSpeedMLP(config, + self.mlp = DeepSpeedMLP(self.config, mp_group, quantize_scales, quantize_groups, @@ -599,7 +622,7 @@ def forward(self, encoder_attention_mask=None, use_cache=False, output_attentions=False): - + #self.config.triangular_masking = False get_present = (get_present or get_key_value or use_cache) input_mask = input_mask if attention_mask is None else attention_mask From 542b25c49ef7e9392ed71da5948665ab3c1ab03a Mon Sep 17 00:00:00 2001 From: Reza Yazdani Date: Wed, 10 Nov 2021 03:24:51 +0500 Subject: [PATCH 2/3] fixing the sparse attention for low block-size --- deepspeed/ops/sparse_attention/matmul.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/deepspeed/ops/sparse_attention/matmul.py b/deepspeed/ops/sparse_attention/matmul.py index 51135bd10579..ea83f093c748 100755 --- a/deepspeed/ops/sparse_attention/matmul.py +++ b/deepspeed/ops/sparse_attention/matmul.py @@ -289,7 +289,7 @@ def make_sdd_lut(layout, block, dtype, device): #_sparse_matmul._load_utils() #start_width = 64 // block #segmented = _sparse_matmul.sdd_segment(layout.type(torch.int32), start_width) - start_width = 128 // block + start_width = (128 if block > 16 else 32) // block layout = layout.type(torch.int32) segmented = libtriton.superblock(layout.data_ptr(), layout.shape[0], From d926bf75b0f1a5cd3c743d8d15a6be4255a1785f Mon Sep 17 00:00:00 2001 From: Reza Yazdani Date: Wed, 10 Nov 2021 03:25:44 +0500 Subject: [PATCH 3/3] remove attn --- .../inference/transformer_inference.py | 30 +------------------ 1 file changed, 1 insertion(+), 29 deletions(-) diff --git a/deepspeed/ops/transformer/inference/transformer_inference.py b/deepspeed/ops/transformer/inference/transformer_inference.py index 7fe4aa5135ad..4f65c121bc8f 100755 --- a/deepspeed/ops/transformer/inference/transformer_inference.py +++ b/deepspeed/ops/transformer/inference/transformer_inference.py @@ -146,31 +146,6 @@ def _transpose_for_context(x): (hidden_size_per_partition,) return x.view(*new_x_layer_shape) - def backup_attention(mixed_query, key_layer, value_layer, input_mask): - if layer_past is not None: - past_key, past_value = layer_past - key_layer = torch.cat((past_key.type_as(key_layer), key_layer), dim=-2) - value_layer = torch.cat((past_value.type_as(value_layer), - value_layer), - dim=-2) - query = _transpose_for_scores(mixed_query, False) - key = _transpose_for_scores(key_layer, True) - value = _transpose_for_scores(value_layer, False) - p = torch.matmul(query, key) - - ds_softmax = inference_cuda_module.softmax_fp16 if config.fp16 else \ - inference_cuda_module.softmax_fp32 - p = ds_softmax(p / (float(key.size(-2))**0.5), - input_mask, - True, - False, - False, - 256) - p = p.to(value.dtype) - context_layer = torch.matmul(p, value) - context_layer = _transpose_for_context(context_layer) - return context_layer, key_layer, value_layer - def compute_attention(qkv_out, input_mask): score_context_func = inference_cuda_module.softmax_context_fp32 if (not config.fp16) else \ inference_cuda_module.softmax_context_fp16 @@ -201,10 +176,7 @@ def compute_attention(qkv_out, input_mask): if no_masking: input_mask = torch.empty(1) head_size = (mixed_query.shape[-1] // num_attention_heads_per_partition) - #return backup_attention(mixed_query, - # key_layer, - # value_layer, - # input_mask) + unfused_mode = not config.specialized_mode or \ mixed_query.shape[1] >= 32 or head_size > 128