fix: resolve audited dispatch, kernel, and rollout bugs

- re-register the linear family with the operator dispatcher (ASTR_OPS / op_backend / resolve)
- fix bf16 gemv misaligned-address faults and element mispairing for offset weights
- reject misaligned bf16_swiglu inputs with a clear error and fall back in the backend gate
- make the rollout reuse decision, validation, and return atomic under one policy snapshot
- add the documented post-scoring rollout version check
- derive live+1 under the scheduler lock in optimizer_step via apply_weight_update(None, ...)
- reject rollout_max_policy_lag below rollout_interval - 1 at config time
- sync gemv stream-test inputs before switching streams; drop dead loader imports
This commit is contained in:
2026-09-03 16:58:14 +08:00
parent 736d1acb2e
commit 7e98a419a7
17 changed files with 485 additions and 49 deletions
+11 -8
View File
@@ -52,15 +52,18 @@ __global__ void bf16_gemv_kernel(
const int wtail_start = whead + wvecs * 8;
const uint4* __restrict__ w4 = reinterpret_cast<const uint4*>(wrow + whead);
// x chunks pair element-for-element with the aligned weight middle. When
// K % 8 == 0 every x row base shares the weight alignment, so one pure
// uint4 loop covers all rows (the production case: head/tail empty, no
// branching inside the loop). Otherwise per-row uint4 loads are not
// 16-byte addressable, and scalar x pairing keeps the kernel correct for
// any K while the weight stream stays vectorized.
// x chunks pair element-for-element with the aligned weight middle:
// the uint4 view is rooted at ``x + whead`` (16-byte aligned by the
// branch guard), and each row strides by ``k / 8`` vectors because its
// first middle element sits ``whead`` scalars past ``row * k``. When
// K % 8 == 0 and the weight row is already aligned (whead == 0, the
// production case) this reduces to one pure uint4 loop with an empty
// head/tail. Otherwise per-row uint4 loads are not 16-byte addressable,
// and scalar x pairing keeps the kernel correct for any K while the
// weight stream stays vectorized.
if (k % 8 == 0 &&
((reinterpret_cast<uintptr_t>(x) + 2u * static_cast<unsigned>(whead)) & 15u) == 0u) {
const auto* x4 = reinterpret_cast<const uint4*>(x);
const auto* x4 = reinterpret_cast<const uint4*>(x + whead);
for (int v = threadIdx.x; v < wvecs; v += blockDim.x) {
const uint4 wv_raw = w4[v];
const auto* wv =
@@ -68,7 +71,7 @@ __global__ void bf16_gemv_kernel(
#pragma unroll
for (int row = 0; row < Rows; ++row) {
const uint4 xv_raw =
x4[(static_cast<int64_t>(row) * wvecs) + v];
x4[(static_cast<int64_t>(row) * (k / 8)) + v];
const auto* xv =
reinterpret_cast<const __nv_bfloat162*>(&xv_raw);
#pragma unroll
+20
View File
@@ -170,6 +170,26 @@ torch::Tensor bf16_swiglu(
gate_weight.is_contiguous(),
"x and weights must be contiguous"
);
// The kernel loads all three streams as uint4; contiguous-but-offset
// views would fault with an opaque "misaligned address" CUDA error, so
// reject them here with an actionable message.
TORCH_CHECK(
(reinterpret_cast<uintptr_t>(x.data_ptr()) & 15u) == 0u,
"bf16_swiglu requires 16-byte aligned x (storage_offset must keep "
"data_ptr divisible by 16); clone the tensor or use the torch path"
);
TORCH_CHECK(
(reinterpret_cast<uintptr_t>(up_weight.data_ptr()) & 15u) == 0u,
"bf16_swiglu requires 16-byte aligned up_weight (storage_offset "
"must keep data_ptr divisible by 16); clone the tensor or use the "
"torch path"
);
TORCH_CHECK(
(reinterpret_cast<uintptr_t>(gate_weight.data_ptr()) & 15u) == 0u,
"bf16_swiglu requires 16-byte aligned gate_weight (storage_offset "
"must keep data_ptr divisible by 16); clone the tensor or use the "
"torch path"
);
TORCH_CHECK(
!x.requires_grad() && !up_weight.requires_grad() &&
!gate_weight.requires_grad(),