File size: 4,196 Bytes
c41750d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
#include "direct_w4a4_m1.h"

#include <torch/extension.h>

#include <optional>

namespace {

torch::Tensor direct_w4a4_m1_linear(
    const torch::Tensor& activation_qdata,
    const torch::Tensor& activation_block_scales,
    const torch::Tensor& activation_tensor_scale,
    const torch::Tensor& weight_qdata,
    const torch::Tensor& weight_block_scales,
    const torch::Tensor& weight_tensor_scale,
    const std::optional<torch::Tensor>& bias) {
  TORCH_CHECK(
      activation_qdata.is_cuda() &&
          activation_qdata.scalar_type() == at::kByte &&
          activation_qdata.dim() == 2 && activation_qdata.is_contiguous(),
      "activation qdata must be contiguous CUDA uint8 [padded_M,K/2]");
  TORCH_CHECK(
      activation_qdata.size(0) >= 1,
      "direct W4A4 M=1 requires at least one packed activation row");
  TORCH_CHECK(
      activation_block_scales.is_cuda() &&
          activation_block_scales.dim() == 2 &&
          activation_block_scales.is_contiguous() &&
          activation_block_scales.element_size() == 1,
      "activation block scales must be contiguous CUDA byte-sized [padded_M,padded_K/16]");
  TORCH_CHECK(
      activation_tensor_scale.is_cuda() &&
          activation_tensor_scale.scalar_type() == at::kFloat &&
          activation_tensor_scale.numel() == 1 &&
          activation_tensor_scale.is_contiguous(),
      "activation tensor scale must be one contiguous CUDA float32 value");
  TORCH_CHECK(
      weight_qdata.is_cuda() && weight_qdata.scalar_type() == at::kByte &&
          weight_qdata.dim() == 2 && weight_qdata.is_contiguous(),
      "weight qdata must be contiguous CUDA uint8 [N,K/2]");
  TORCH_CHECK(
      weight_block_scales.is_cuda() && weight_block_scales.dim() == 2 &&
          weight_block_scales.is_contiguous() &&
          weight_block_scales.element_size() == 1,
      "weight block scales must be contiguous CUDA byte-sized [padded_N,padded_K/16]");
  TORCH_CHECK(
      weight_tensor_scale.is_cuda() &&
          weight_tensor_scale.scalar_type() == at::kFloat &&
          weight_tensor_scale.numel() == 1 &&
          weight_tensor_scale.is_contiguous(),
      "weight tensor scale must be one contiguous CUDA float32 value");
  TORCH_CHECK(
      activation_qdata.device() == activation_block_scales.device() &&
          activation_qdata.device() == activation_tensor_scale.device() &&
          activation_qdata.device() == weight_qdata.device() &&
          activation_qdata.device() == weight_block_scales.device() &&
          activation_qdata.device() == weight_tensor_scale.device(),
      "all direct W4A4 tensors must use the same CUDA device");

  const int64_t in_features = activation_qdata.size(1) * 2;
  const int64_t out_features = weight_qdata.size(0);
  TORCH_CHECK(
      in_features > 0 && in_features % 32 == 0,
      "direct W4A4 M=1 requires K divisible by 32");
  TORCH_CHECK(
      weight_qdata.size(1) == activation_qdata.size(1),
      "activation and weight packed K dimensions differ");
  TORCH_CHECK(out_features > 0, "direct W4A4 M=1 requires positive N");
  TORCH_CHECK(
      activation_block_scales.size(0) >= 1 &&
          activation_block_scales.size(1) >= in_features / 16,
      "activation block-scale tensor is too small");
  TORCH_CHECK(
      weight_block_scales.size(0) >= out_features &&
          weight_block_scales.size(1) >= in_features / 16,
      "weight block-scale tensor is too small");

  if (bias.has_value()) {
    const auto& value = *bias;
    TORCH_CHECK(
        value.is_cuda() && value.scalar_type() == at::kBFloat16 &&
            value.dim() == 1 && value.is_contiguous() &&
            value.numel() == out_features &&
            value.device() == activation_qdata.device(),
        "bias must be contiguous CUDA bfloat16 [N]");
  }
  return direct_w4a4_m1_linear_cuda(
      activation_qdata,
      activation_block_scales,
      activation_tensor_scale,
      weight_qdata,
      weight_block_scales,
      weight_tensor_scale,
      bias);
}

}  // namespace

PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
  module.def(
      "linear",
      &direct_w4a4_m1_linear,
      "Direct packed-NVFP4 activation x packed-NVFP4 weight M=1 linear");
}