Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 15 additions & 0 deletions python/freetoken/engine/engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -1166,6 +1166,21 @@ def _cpu_moe_executor_viable(model_config) -> bool:
return False
expert_quant = getattr(model_config, "expert_quant", "none")
fmt = expert_quant if expert_quant != "none" else (moe_wfmt or "bf16")
if fmt == "gguf":
# "gguf" is a container tag, not a layout: the checkpoint picks a ggml type per
# tensor and the concrete CPU format has to be recovered from the bank types.
# Testing the tag against _WFMT_IDS answers False for EVERY GGUF checkpoint, which
# silently disables the automatic residency split on hosts where CUDA pinning is
# quota-capped (WSL caps it near half of RAM). The symptom is not a clear refusal
# but cudaHostRegister failing partway through the banks.
from freetoken.moe.cpu_executor import _GGML_TO_CPU_FMT

types = getattr(model_config, "gguf_expert_types", None)
if not types:
return False
gate_up, down = int(types[0]), int(types[1])
# one weight_format serves both banks, so mixed types cannot run on the CPU path
return gate_up == down and gate_up in _GGML_TO_CPU_FMT
return fmt == "mxfp4" or fmt in _WFMT_IDS


Expand Down
7 changes: 7 additions & 0 deletions python/freetoken/kernel/csrc/cpu_moe/cpu_moe_ext.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1650,6 +1650,13 @@ struct CpuMoeExecutor {
const uint8_t* w = gu_packed_l + ((size_t)e * (2 * I) + row) * (size_t)q4_gu_row_bytes;
return q6_k_dot_f32_scalar(w, x, H); // W4A16: bf16 activations, K-quant dequant
}
// Anything that reaches here is assumed NVFP4 and dereferences the scale/global
// pointers, which are null for formats that do not have them (the GGUF banks pass 0).
// Falling through with an unhandled format therefore segfaults inside the worker
// thread rather than reporting anything useful, so reject it here instead.
TORCH_CHECK(fmt == WF_NVFP4 || fmt == WF_DSFP4,
"cpu_moe gemm1_dot: unhandled weight_format ", fmt,
" (handled: bf16=0, nvfp4=1, mxfp4=2, dsfp4=3, q4_0=4, q4_k=5, q6_k=6)");
const size_t r = (size_t)e * (2 * I) + row;
if (use_vnni)
return nvi8dot(gu_packed_l + r * (size_t)(H / 2), gu_scale_l + r * (size_t)(H / 16),
Expand Down
40 changes: 40 additions & 0 deletions python/freetoken/kernel/csrc/gguf/gguf_dequant_kernel.cu
Original file line number Diff line number Diff line change
@@ -0,0 +1,40 @@
// ROCm operation-split binding for the vendored GGUF dequant kernels.
// Kernel implementations remain in the sgl-kernel/llama.cpp-derived headers.
#include <c10/cuda/CUDAGuard.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <torch/all.h>

#include "dispatch.h"
#include "ggml-common.h"
#include "dequantize.cuh"

torch::Tensor ggml_dequantize(
torch::Tensor W,
int64_t type,
int64_t m,
int64_t n,
std::optional<at::ScalarType> const& dtype) {
const at::cuda::OptionalCUDAGuard device_guard(device_of(W));
auto dtype_ = dtype.value_or(torch::kFloat16);
auto options = torch::TensorOptions().dtype(dtype_).device(W.device());
at::Tensor DW = torch::empty({m, n}, options);
cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();

DISPATCH_FLOAT_TYPES(DW.scalar_type(), "ggml_dequantize", [&] {
auto to_cuda = ggml_get_to_cuda<scalar_t>(type);
TORCH_CHECK(
to_cuda != nullptr,
"ggml_dequantize: unsupported GGUF quant type ", type,
" (dequant kernels exist for Q4_0/Q4_1/Q5_0/Q5_1/Q8_0/Q2_K-Q6_K/IQ2_XXS/"
"IQ2_XS/IQ3_XXS/IQ1_S/IQ4_NL/IQ3_S/IQ2_S/IQ4_XS/IQ1_M)");
to_cuda((void*)W.data_ptr(), (scalar_t*)DW.data_ptr(), m * n, stream);
});

return DW;
}

#include <torch/extension.h>
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("ggml_dequantize", &ggml_dequantize, "");
}
57 changes: 57 additions & 0 deletions python/freetoken/kernel/csrc/gguf/gguf_mmq_kernel.cu
Original file line number Diff line number Diff line change
@@ -0,0 +1,57 @@
// ROCm operation-split binding for GGUF large-batch MMQ.
#include <c10/cuda/CUDAGuard.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <torch/all.h>

#include "dispatch.h"
#include "ggml-common.h"
#include "vecdotq.cuh"
#include "mmq.cuh"
#include "quantize_q8_1.cuh"

torch::Tensor ggml_mul_mat_a8(
torch::Tensor W,
torch::Tensor X,
int64_t type,
int64_t row) {
int col = X.sizes()[1];
int padded = (col + 512 - 1) / 512 * 512;
int batch = X.sizes()[0];
const at::cuda::OptionalCUDAGuard device_guard(device_of(X));
auto options = torch::TensorOptions().dtype(X.dtype()).device(W.device());
at::Tensor Y = torch::empty({batch, row}, options);
cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
options = torch::TensorOptions().dtype(torch::kInt32).device(W.device());
at::Tensor quant_X = torch::empty({batch, padded / 32 * 9}, options);
DISPATCH_FLOAT_TYPES(X.scalar_type(), "ggml_mul_mat_a8", [&] {
quantize_row_q8_1_cuda<scalar_t>(
(scalar_t*)X.data_ptr(), (void*)quant_X.data_ptr(), col, batch, stream);
using Fn = void (*)(const void*, const void*, scalar_t*, int, int, int, int, int, cudaStream_t);
Fn fn = nullptr;
switch (type) {
case 2: fn = &ggml_mul_mat_q4_0_q8_1_cuda<scalar_t>; break;
case 3: fn = &ggml_mul_mat_q4_1_q8_1_cuda<scalar_t>; break;
case 6: fn = &ggml_mul_mat_q5_0_q8_1_cuda<scalar_t>; break;
case 7: fn = &ggml_mul_mat_q5_1_q8_1_cuda<scalar_t>; break;
case 8: fn = &ggml_mul_mat_q8_0_q8_1_cuda<scalar_t>; break;
case 10: fn = &ggml_mul_mat_q2_K_q8_1_cuda<scalar_t>; break;
case 11: fn = &ggml_mul_mat_q3_K_q8_1_cuda<scalar_t>; break;
case 12: fn = &ggml_mul_mat_q4_K_q8_1_cuda<scalar_t>; break;
case 13: fn = &ggml_mul_mat_q5_K_q8_1_cuda<scalar_t>; break;
case 14: fn = &ggml_mul_mat_q6_K_q8_1_cuda<scalar_t>; break;
default:
TORCH_CHECK(false, "ggml_mul_mat_a8: unsupported GGUF quant type ", type,
" (MMQ kernels exist only for Q4_0/Q4_1/Q5_0/Q5_1/Q8_0/Q2_K-Q6_K; "
"I-quants must route through ggml_dequantize)");
}
fn(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(),
col, row, batch, padded, row, stream);
});
return Y;
}

#include <torch/extension.h>
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("ggml_mul_mat_a8", &ggml_mul_mat_a8, "");
}
102 changes: 102 additions & 0 deletions python/freetoken/kernel/csrc/gguf/gguf_mmvq_kernel.cu
Original file line number Diff line number Diff line change
@@ -0,0 +1,102 @@
// ROCm operation-split binding for GGUF small-batch MMVQ.
#include <c10/cuda/CUDAGuard.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <torch/all.h>

#include "dispatch.h"
#include "ggml-common.h"
#include "vecdotq.cuh"
#include "mmvq.cuh"
#include "quantize_q8_1.cuh"

torch::Tensor ggml_mul_mat_vec_a8(
torch::Tensor W,
torch::Tensor X,
int64_t type,
int64_t row) {
int col = X.sizes()[1];
int vecs = X.sizes()[0];
const int padded = (col + 512 - 1) / 512 * 512;
const at::cuda::OptionalCUDAGuard device_guard(device_of(X));
auto options = torch::TensorOptions().dtype(X.dtype()).device(W.device());
at::Tensor Y = torch::empty({vecs, row}, options);
cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
options = torch::TensorOptions().dtype(torch::kInt32).device(W.device());
at::Tensor quant_X = torch::empty({vecs, padded / 32 * 9}, options);
DISPATCH_FLOAT_TYPES(X.scalar_type(), "ggml_mul_mat_vec_a8", [&] {
quantize_row_q8_1_cuda<scalar_t>(
(scalar_t*)X.data_ptr(), (void*)quant_X.data_ptr(), col, vecs, stream);
switch (type) {
case 2:
mul_mat_vec_q4_0_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 3:
mul_mat_vec_q4_1_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 6:
mul_mat_vec_q5_0_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 7:
mul_mat_vec_q5_1_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 8:
mul_mat_vec_q8_0_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 10:
mul_mat_vec_q2_K_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 11:
mul_mat_vec_q3_K_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 12:
mul_mat_vec_q4_K_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 13:
mul_mat_vec_q5_K_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 14:
mul_mat_vec_q6_K_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 16:
mul_mat_vec_iq2_xxs_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 17:
mul_mat_vec_iq2_xs_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 18:
mul_mat_vec_iq3_xxs_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 19:
mul_mat_vec_iq1_s_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 20:
mul_mat_vec_iq4_nl_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 21:
mul_mat_vec_iq3_s_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 22:
mul_mat_vec_iq2_s_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 23:
mul_mat_vec_iq4_xs_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
case 29:
mul_mat_vec_iq1_m_q8_1_cuda<scalar_t>(W.data_ptr(), quant_X.data_ptr(), (scalar_t*)Y.data_ptr(), col, row, vecs, stream);
break;
default:
TORCH_CHECK(
false,
"ggml_mul_mat_vec_a8: unsupported GGUF quant type ", type,
" (MMVQ kernels exist for Q4_0/Q4_1/Q5_0/Q5_1/Q8_0/Q2_K-Q6_K/IQ2_XXS/IQ2_XS/"
"IQ3_XXS/IQ1_S/IQ4_NL/IQ3_S/IQ2_S/IQ4_XS/IQ1_M)");
}
});
return Y;
}

#include <torch/extension.h>
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("ggml_mul_mat_vec_a8", &ggml_mul_mat_vec_a8, "");
}
87 changes: 87 additions & 0 deletions python/freetoken/kernel/csrc/gguf/gguf_moe_kernel.cu
Original file line number Diff line number Diff line change
@@ -0,0 +1,87 @@
// ROCm operation-split binding for GGUF grouped large-batch MoE kernels.
#include <c10/cuda/CUDAGuard.h>
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <torch/all.h>

#include "dispatch.h"
#include "ggml-common.h"
#include "vecdotq.cuh"
#include "mmq.cuh"
#include "moe.cuh"
#include "quantize_q8_1.cuh"

torch::Tensor ggml_moe_a8(
torch::Tensor X,
torch::Tensor W,
torch::Tensor sorted_token_ids,
torch::Tensor expert_ids,
torch::Tensor num_tokens_post_padded,
int64_t type,
int64_t row,
int64_t top_k,
int64_t tokens) {
int col = X.sizes()[1];
int padded = (col + 512 - 1) / 512 * 512;
const at::cuda::OptionalCUDAGuard device_guard(device_of(X));
auto options = torch::TensorOptions().dtype(X.dtype()).device(W.device());
at::Tensor Y = torch::empty({tokens * top_k, row}, options);
cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
options = torch::TensorOptions().dtype(torch::kInt32).device(W.device());
at::Tensor quant_X = torch::empty({tokens, padded / 32 * 9}, options);
DISPATCH_FLOAT_TYPES(X.scalar_type(), "ggml_moe_a8", [&] {
quantize_row_q8_1_cuda<scalar_t>(
(scalar_t*)X.data_ptr(), (void*)quant_X.data_ptr(), col, tokens, stream);
using Fn = void (*)(
const void*, const void*, scalar_t*, const int*, const int*, const int*,
int, int, int, int, int, int, int, int, cudaStream_t);
Fn fn = nullptr;
switch (type) {
case 2: fn = &ggml_moe_q4_0_q8_1_cuda<scalar_t>; break;
case 3: fn = &ggml_moe_q4_1_q8_1_cuda<scalar_t>; break;
case 6: fn = &ggml_moe_q5_0_q8_1_cuda<scalar_t>; break;
case 7: fn = &ggml_moe_q5_1_q8_1_cuda<scalar_t>; break;
case 8: fn = &ggml_moe_q8_0_q8_1_cuda<scalar_t>; break;
case 10: fn = &ggml_moe_q2_K_q8_1_cuda<scalar_t>; break;
case 11: fn = &ggml_moe_q3_K_q8_1_cuda<scalar_t>; break;
case 12: fn = &ggml_moe_q4_K_q8_1_cuda<scalar_t>; break;
case 13: fn = &ggml_moe_q5_K_q8_1_cuda<scalar_t>; break;
case 14: fn = &ggml_moe_q6_K_q8_1_cuda<scalar_t>; break;
default:
TORCH_CHECK(false, "ggml_moe_a8: unsupported GGUF quant type ", type,
" (MMQ kernels exist only for Q4_0/Q4_1/Q5_0/Q5_1/Q8_0/Q2_K-Q6_K; "
"I-quants must route through ggml_dequantize)");
}
fn(quant_X.data_ptr(), W.data_ptr(), (scalar_t*)Y.data_ptr(),
(int*)sorted_token_ids.data_ptr(), (int*)expert_ids.data_ptr(),
(int*)num_tokens_post_padded.data_ptr(), W.stride(0), col, row, tokens,
padded, row, top_k, sorted_token_ids.sizes()[0], stream);
});
return Y;
}

int64_t ggml_moe_get_block_size(int64_t type) {
switch (type) {
case 2: return MOE_X_Q4_0;
case 3: return MOE_X_Q4_1;
case 6: return MOE_X_Q5_0;
case 7: return MOE_X_Q5_1;
case 8: return MOE_X_Q8_0;
case 10: return MOE_X_Q2_K;
case 11: return MOE_X_Q3_K;
case 12: return MOE_X_Q4_K;
case 13: return MOE_X_Q5_K;
case 14: return MOE_X_Q6_K;
default:
TORCH_CHECK(false, "ggml_moe_get_block_size: unsupported GGUF quant type ", type,
" (MMQ kernels exist only for Q4_0/Q4_1/Q5_0/Q5_1/Q8_0/Q2_K-Q6_K; "
"I-quants must route through ggml_dequantize)");
return 0;
}
}

#include <torch/extension.h>
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("ggml_moe_a8", &ggml_moe_a8, "");
m.def("ggml_moe_get_block_size", &ggml_moe_get_block_size, "");
}
Loading