Follow up to #27542
(1) If we do "temp = Concat(past_key, new_key)", then this Concat will copy the past_key to temp. But we end up doing a second copy when we do a mem-cpy from temp to present_key. If we "present_key = Concat(past_key, new_key)", we will avoid the second copy (while we still do the first copy in Concat itself).
|
// Copy past KV (BNSH) into present buffers (BNSH). Strided copy because |
|
// past has [B, N_kv, past_seq, H] and present has [B, N_kv, total_seq, H]. |
|
const size_t past_k_row_bytes = static_cast<size_t>(parameters.past_sequence_length) * |
|
parameters.head_size * sizeof(T); |
|
const size_t present_k_row_bytes = static_cast<size_t>(parameters.total_sequence_length) * |
|
parameters.head_size * sizeof(T); |
|
CUDA_RETURN_IF_ERROR(cudaMemcpy2DAsync( |
|
present_key->MutableData<T>(), present_k_row_bytes, |
|
past_key->Data<T>(), past_k_row_bytes, |
|
past_k_row_bytes, num_kv_rows, |
|
cudaMemcpyDeviceToDevice, cuda_stream)); |
|
|
|
const size_t past_v_row_bytes = static_cast<size_t>(parameters.past_sequence_length) * |
|
parameters.v_head_size * sizeof(T); |
|
const size_t present_v_row_bytes = static_cast<size_t>(parameters.total_sequence_length) * |
|
parameters.v_head_size * sizeof(T); |
|
CUDA_RETURN_IF_ERROR(cudaMemcpy2DAsync( |
|
present_value->MutableData<T>(), present_v_row_bytes, |
|
past_value->Data<T>(), past_v_row_bytes, |
|
past_v_row_bytes, num_kv_rows, |
|
cudaMemcpyDeviceToDevice, cuda_stream)); |
(2) CUDA_KERNEL_ASSERT is noop in release build. Check the convention in onnxruntime to see an official way to raise for invalid.
|
CUDA_KERNEL_ASSERT(false); // mask must be contiguous (no True after False) |
#27542 (comment)
Follow up to #27542
(1) If we do "temp = Concat(past_key, new_key)", then this Concat will copy the past_key to temp. But we end up doing a second copy when we do a mem-cpy from temp to present_key. If we "present_key = Concat(past_key, new_key)", we will avoid the second copy (while we still do the first copy in Concat itself).
onnxruntime/onnxruntime/core/providers/cuda/llm/attention.cc
Lines 327 to 347 in 56bce08
(2) CUDA_KERNEL_ASSERT is noop in release build. Check the convention in onnxruntime to see an official way to raise for invalid.
onnxruntime/onnxruntime/core/providers/cuda/llm/attention_mask_impl.cu
Line 96 in 56bce08
#27542 (comment)