mirror of
https://github.com/ROCm/composable_kernel.git
synced 2026-05-02 04:31:25 +00:00
[CK] Add FP8 KV_BLOCKSCALE support for batch prefill MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Implement per-page K/V quantization for paged attention: - Add KV_BLOCKSCALE enum to BlockAttentionQuantScaleEnum - Use exp2 shift trick to eliminate explicit P scaling overhead - Prefetch physical pages offset for KV cache, overlaps with computations ## Proposed changes Please describe the motivation behind the pull request, whether it enables a new feature or fixes a bug. If there are associated pull requests or issues, please link them to the pull request. ## Checklist Please put an `x` into the boxes that apply. You can also fill these out after creating the PR. If you're not sure, please don't hesitate to ask. - [ ] I have added tests relevant to the introduced functionality, and the unit tests are passing locally - [ ] I have added the test to REGRESSION_TESTS list defined at the top of CMakeLists.txt in tests/CMakeLists.txt, **IF** the test takes more than 30 seconds to run. - [ ] I have added inline documentation which enables the maintainers with understanding the motivation - [ ] I have removed the stale documentation which is no longer relevant after this pull request - [ ] (If this change is user-facing) I have added release notes which provide the end users with a brief summary of the improvement from this pull request - [ ] I have run `clang-format` on all changed files - [ ] Any dependent changes have been merged ## Discussion If this is a relatively large or complex change, feel free to start a discussion by explaining why you chose the solution you did and what alternatives you considered
72 lines
1.8 KiB
C++
72 lines
1.8 KiB
C++
// Copyright (c) Advanced Micro Devices, Inc., or its affiliates.
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
#pragma once
|
|
|
|
#include <ostream>
|
|
#include <string>
|
|
#include "ck_tile/core.hpp"
|
|
#include "ck_tile/ops/fmha.hpp"
|
|
|
|
#pragma clang diagnostic push
|
|
#pragma clang diagnostic ignored "-Wlifetime-safety-intra-tu-suggestions"
|
|
|
|
// keep sync with BlockAttentionQuantScaleEnum
|
|
enum class quant_scale_enum
|
|
{
|
|
no_scale = 0,
|
|
pertensor = 1,
|
|
blockscale = 2,
|
|
kv_blockscale = 3, // Q per-tensor, K/V per-page block scale
|
|
};
|
|
|
|
struct quant_scale_info
|
|
{
|
|
quant_scale_enum type;
|
|
|
|
void serialize(std::ostream& os) const
|
|
{
|
|
if(type == quant_scale_enum::no_scale)
|
|
os << "n";
|
|
else if(type == quant_scale_enum::pertensor)
|
|
os << "pt";
|
|
else if(type == quant_scale_enum::blockscale)
|
|
os << "bs";
|
|
else if(type == quant_scale_enum::kv_blockscale)
|
|
os << "kvbs";
|
|
}
|
|
|
|
static quant_scale_info decode(std::string str)
|
|
{
|
|
quant_scale_info info{quant_scale_enum::no_scale};
|
|
if(str == "n" || str == "0")
|
|
{
|
|
info.type = quant_scale_enum::no_scale;
|
|
}
|
|
else if(str == "pt" || str == "1")
|
|
{
|
|
info.type = quant_scale_enum::pertensor;
|
|
}
|
|
else if(str == "bs" || str == "2")
|
|
{
|
|
info.type = quant_scale_enum::blockscale;
|
|
}
|
|
else if(str == "kvbs" || str == "3")
|
|
{
|
|
info.type = quant_scale_enum::kv_blockscale;
|
|
}
|
|
else
|
|
{
|
|
throw std::invalid_argument("invalid quant scale value: " + str);
|
|
}
|
|
return info;
|
|
}
|
|
|
|
friend std::ostream& operator<<(std::ostream& os, const quant_scale_info& qsi)
|
|
{
|
|
qsi.serialize(os);
|
|
return os;
|
|
}
|
|
};
|
|
#pragma clang diagnostic pop
|