v4.1 release
This commit is contained in:
@@ -75,15 +75,7 @@ __global__ void dequantize_kernel(DequantizedElement* dq_buffer,
|
||||
// Represent the full tensors to gmem elements.
|
||||
// These are expected to have shape [MN, K, L]
|
||||
cute::Tensor gmem_op_dq = cute::make_tensor(cute::make_gmem_ptr(dq_buffer), operand_layout);
|
||||
auto init_quantized_iterator = [&]() {
|
||||
if constexpr (cute::sizeof_bits_v<QuantizedElement> >= 8) {
|
||||
return cute::make_gmem_ptr(q_buffer);
|
||||
}
|
||||
else {
|
||||
return cute::subbyte_iterator<const QuantizedElement>(q_buffer);
|
||||
}
|
||||
};
|
||||
cute::Tensor gmem_op_q = cute::make_tensor(init_quantized_iterator(), operand_layout);
|
||||
cute::Tensor gmem_op_q = cute::make_tensor(cute::make_gmem_ptr<QuantizedElement const>(q_buffer), operand_layout);
|
||||
// While the scales are expected to have shape [MN, G, L] but with a stride to allow broadcasting
|
||||
// It is expected that K % G == 0
|
||||
cute::Tensor gmem_scale_broadcasted = cute::make_tensor(make_gmem_ptr(scale_buffer), broadcasted_scale_layout);
|
||||
|
||||
@@ -95,6 +95,7 @@ template<
|
||||
class ElementCompute_,
|
||||
class ElementC_,
|
||||
class ElementOut_,
|
||||
bool ResidualAdd_,
|
||||
class TensorAlpha_,
|
||||
class TensorBeta_,
|
||||
class TensorBias_,
|
||||
@@ -110,6 +111,8 @@ struct ConvEpilogueFusionParams {
|
||||
using TensorBeta = TensorBeta_;
|
||||
using TensorBias = TensorBias_;
|
||||
using ActivationFunctor = ActivationFunctor_;
|
||||
static constexpr bool ResidualAdd = ResidualAdd_; // Source added after activation
|
||||
|
||||
ElementScalar alpha = ElementScalar(1);
|
||||
ElementScalar beta = ElementScalar(0);
|
||||
|
||||
@@ -228,12 +231,17 @@ private:
|
||||
epi_fusion_params_.tensor_alpha[k] : epi_fusion_params_.alpha;
|
||||
ElementScalar beta = raw_pointer_cast(epi_fusion_params_.tensor_beta.data()) ?
|
||||
epi_fusion_params_.tensor_beta[k] : epi_fusion_params_.beta;
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator) +
|
||||
scale_converter(beta) * residual_converter(tensor_c_(k, q, n, g));
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator);
|
||||
if (not EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(k, q, n, g));
|
||||
}
|
||||
if (raw_pointer_cast(epi_fusion_params_.tensor_bias.data())) {
|
||||
output += bias_converter(epi_fusion_params_.tensor_bias[k]);
|
||||
}
|
||||
output = epi_activation(output);
|
||||
if (EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(k, q, n, g));
|
||||
}
|
||||
tensor_d_(k, q, n, g) = output_converter(output);
|
||||
}
|
||||
}
|
||||
@@ -279,12 +287,17 @@ private:
|
||||
epi_fusion_params_.tensor_alpha[k] : epi_fusion_params_.alpha;
|
||||
ElementScalar beta = raw_pointer_cast(epi_fusion_params_.tensor_beta.data()) ?
|
||||
epi_fusion_params_.tensor_beta[k] : epi_fusion_params_.beta;
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator) +
|
||||
scale_converter(beta) * residual_converter(tensor_c_(k, q, p, n, g));
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator);
|
||||
if (not EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(k, q, p, n, g));
|
||||
}
|
||||
if (raw_pointer_cast(epi_fusion_params_.tensor_bias.data())) {
|
||||
output += bias_converter(epi_fusion_params_.tensor_bias[k]);
|
||||
}
|
||||
output = epi_activation(output);
|
||||
if (EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(k, q, p, n, g));
|
||||
}
|
||||
tensor_d_(k, q, p, n, g) = output_converter(output);
|
||||
}
|
||||
}
|
||||
@@ -337,12 +350,17 @@ private:
|
||||
epi_fusion_params_.tensor_alpha[k] : epi_fusion_params_.alpha;
|
||||
ElementScalar beta = raw_pointer_cast(epi_fusion_params_.tensor_beta.data()) ?
|
||||
epi_fusion_params_.tensor_beta[k] : epi_fusion_params_.beta;
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator) +
|
||||
scale_converter(beta) * residual_converter(tensor_c_(k, q, p, z, n, g));
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator);
|
||||
if (not EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(k, q, p, z, n, g));
|
||||
}
|
||||
if (raw_pointer_cast(epi_fusion_params_.tensor_bias.data())) {
|
||||
output += bias_converter(epi_fusion_params_.tensor_bias[k]);
|
||||
}
|
||||
output = epi_activation(output);
|
||||
if (EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(k, q, p, z, n, g));
|
||||
}
|
||||
tensor_d_(k, q, p, z, n, g) = output_converter(output);
|
||||
}
|
||||
}
|
||||
@@ -389,12 +407,17 @@ private:
|
||||
? epi_fusion_params_.tensor_alpha[c] : epi_fusion_params_.alpha;
|
||||
ElementScalar beta = raw_pointer_cast(epi_fusion_params_.tensor_beta.data())
|
||||
? epi_fusion_params_.tensor_beta[c] : epi_fusion_params_.beta;
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator) +
|
||||
scale_converter(beta) * residual_converter(tensor_c_(c, w, n, g));
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator);
|
||||
if (not EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(c, w, n, g));
|
||||
}
|
||||
if (raw_pointer_cast(epi_fusion_params_.tensor_bias.data())) {
|
||||
output += bias_converter(epi_fusion_params_.tensor_bias[c]);
|
||||
}
|
||||
output = epi_activation(output);
|
||||
if (EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(c, w, n, g));
|
||||
}
|
||||
tensor_d_(c, w, n, g) = output_converter(output);
|
||||
}
|
||||
}
|
||||
@@ -451,12 +474,17 @@ private:
|
||||
? epi_fusion_params_.tensor_alpha[c] : epi_fusion_params_.alpha;
|
||||
ElementScalar beta = raw_pointer_cast(epi_fusion_params_.tensor_beta.data())
|
||||
? epi_fusion_params_.tensor_beta[c] : epi_fusion_params_.beta;
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator) +
|
||||
scale_converter(beta) * residual_converter(tensor_c_(c, w, h, n, g));
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator);
|
||||
if (not EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(c, w, h, n, g));
|
||||
}
|
||||
if (raw_pointer_cast(epi_fusion_params_.tensor_bias.data())) {
|
||||
output += bias_converter(epi_fusion_params_.tensor_bias[c]);
|
||||
}
|
||||
output = epi_activation(output);
|
||||
if (EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(c, w, h, n, g));
|
||||
}
|
||||
|
||||
tensor_d_(c, w, h, n, g) = output_converter(output);
|
||||
}
|
||||
@@ -527,12 +555,17 @@ private:
|
||||
? epi_fusion_params_.tensor_alpha[c] : epi_fusion_params_.alpha;
|
||||
ElementScalar beta = raw_pointer_cast(epi_fusion_params_.tensor_beta.data())
|
||||
? epi_fusion_params_.tensor_beta[c] : epi_fusion_params_.beta;
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator) +
|
||||
scale_converter(beta) * residual_converter(tensor_c_(c, w, h, d, n, g));
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator);
|
||||
if (not EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(c, w, h, d, n, g));
|
||||
}
|
||||
if (raw_pointer_cast(epi_fusion_params_.tensor_bias.data())) {
|
||||
output += bias_converter(epi_fusion_params_.tensor_bias[c]);
|
||||
}
|
||||
output = epi_activation(output);
|
||||
if (EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(c, w, h, d, n, g));
|
||||
}
|
||||
tensor_d_(c, w, h, d, n, g) = output_converter(output);
|
||||
}
|
||||
}
|
||||
@@ -583,12 +616,17 @@ private:
|
||||
ElementScalar beta = raw_pointer_cast(epi_fusion_params_.tensor_beta.data()) ?
|
||||
epi_fusion_params_.tensor_beta[c] : epi_fusion_params_.beta;
|
||||
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator) +
|
||||
scale_converter(beta) * residual_converter(tensor_c_(c, s, k, g));
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator);
|
||||
if (not EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(c, s, k, g));
|
||||
}
|
||||
if (raw_pointer_cast(epi_fusion_params_.tensor_bias.data())) {
|
||||
output += bias_converter(epi_fusion_params_.tensor_bias[c]);
|
||||
}
|
||||
output = epi_activation(output);
|
||||
if (EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(c, s, k, g));
|
||||
}
|
||||
tensor_d_(c, s, k, g) = output_converter(output);
|
||||
}
|
||||
}
|
||||
@@ -643,12 +681,17 @@ private:
|
||||
ElementScalar beta = raw_pointer_cast(epi_fusion_params_.tensor_beta.data()) ?
|
||||
epi_fusion_params_.tensor_beta[c] : epi_fusion_params_.beta;
|
||||
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator) +
|
||||
scale_converter(beta) * residual_converter(tensor_c_(c, s, r, k, g));
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator);
|
||||
if (not EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(c, s, r, k, g));
|
||||
}
|
||||
if (raw_pointer_cast(epi_fusion_params_.tensor_bias.data())) {
|
||||
output += bias_converter(epi_fusion_params_.tensor_bias[c]);
|
||||
}
|
||||
output = epi_activation(output);
|
||||
if (EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(c, s, r, k, g));
|
||||
}
|
||||
tensor_d_(c, s, r, k, g) = output_converter(output);
|
||||
}
|
||||
}
|
||||
@@ -711,12 +754,17 @@ private:
|
||||
ElementScalar beta = raw_pointer_cast(epi_fusion_params_.tensor_beta.data()) ?
|
||||
epi_fusion_params_.tensor_beta[c] : epi_fusion_params_.beta;
|
||||
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator) +
|
||||
scale_converter(beta) * residual_converter(tensor_c_(c, s, r, t, k, g));
|
||||
ElementCompute output = scale_converter(alpha) * acc_converter(accumulator);
|
||||
if (not EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(c, s, r, t, k, g));
|
||||
}
|
||||
if (raw_pointer_cast(epi_fusion_params_.tensor_bias.data())) {
|
||||
output += bias_converter(epi_fusion_params_.tensor_bias[c]);
|
||||
}
|
||||
output = epi_activation(output);
|
||||
if (EpilogueFusionParams::ResidualAdd) {
|
||||
output += scale_converter(beta) * residual_converter(tensor_c_(c, s, r, t, k, g));
|
||||
}
|
||||
tensor_d_(c, s, r, t, k, g) = output_converter(output);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user