Skip to content
Open
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
32 changes: 30 additions & 2 deletions python/quadrants/lang/_func_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -83,6 +83,20 @@ def _kernel_coverage_enabled() -> bool:
_arch_cuda = _qd_core.Arch.cuda
_is_cpython = sys.implementation.name == "cpython"

# The ndarray ExternalPtrStmt linear index is accumulated in int64 in the LLVM codegen backends
# (see TaskCodeGenLLVM::visit(ExternalPtrStmt)) and, on the SPIR-V backends (Vulkan/Metal/...), whenever the
# device advertises 64-bit integers (DeviceCapability.spirv_has_int64 -> TaskCodegen::visit widens to i64).
# Only a SPIR-V device without shaderInt64 still flattens the offset in int32 and can overflow, so the
# product-of-dimensions warning below is gated on both the arch and the device capability.
_ARCHS_WITH_I64_LINEAR_INDEX = frozenset(
{
_qd_core.Arch.x64,
_qd_core.Arch.arm64,
_qd_core.Arch.cuda,
_qd_core.Arch.amdgpu,
}
)

# PERF: Frozen-dataclass dispatch caching.
#
# When a frozen dataclass (e.g. Genesis's StructConstraintState with ~43 fields) is passed to a kernel, the per-launch
Expand Down Expand Up @@ -727,14 +741,28 @@ def _recursive_set_args(
# array shapes.
is_soa = needed_arg_type.layout == Layout.SOA
array_shape = v.shape
if math.prod(array_shape) > np.iinfo(np.int32).max:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

wait, there was already a warning?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes — main had an unconditional one here: warnings.warn("...int64 indexing is not supported yet."). But it was only a warning (the offset still overflowed and returned the wrong element), and its message is now false for CPU/CUDA/AMDGPU since those accumulate in int64. Keeping it as-is would warn on exactly the backends this PR fixes. So I kept a warning but scoped it to the backends that still flatten in int32 — and with the new SPIR-V int64 commit, only to devices without shaderInt64.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I would prefer to just keep the warning as-is, without changing anything, without a strong reason.

What is your motivation for creating this PR? I'm sensing the motivation is primarily 'this is broken and it should be fixed'?

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

(I'm happy to upgrade from warning to throwing an exception, if that would make you more comfortable?)

@paveltc paveltc Jul 18, 2026

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Here is the motivation from a colleague of mine who worked on this:

Things that I'm wary about:

  • changes without a clear concrete benefit, since any change involves risk, of adding new bugs, and of performance regressions
    They benefit was that they enabled workloads with more than 32k envs (which we explicitly needed)

  • performance regressions, since 64-bit types tend to be slower than 32-bit types, e.g. because use more memory bandwidth
    Sure 64 bit types are slow, but we explicitly don't switch to using int64 for indexes, we narrowed the scope and used it for things like flattened indices which show up a lot and can easily exceed the bounds of int32

  • correctness regressions, and crash bugs, which changing bit-widths can definitely both create, especially since we'd need to change this across all platforms
    The hope was that the limited scope of these changes would also reduce or eliminate these types of problems.

@paveltc paveltc Jul 18, 2026

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Quick update, and a correction to my earlier comment. When I said genesis-world doesn't hit this — that physics slows down before an ndarray crosses ~2.1B elements — I was reasoning from my own assumptions. Since then a colleague who worked on the original change filled me in: they did hit it in practice, on a real workload running >32k parallel envs, where the aggregate ndarray crosses INT32_MAX even though the per-env tensors are modest. So there is a concrete, non-artificial workload behind this, not just the memory-safety argument — I was wrong to imply nobody needed the capability.

That also speaks to the warning-vs-fix question. For that workload, neither the existing warning nor upgrading it to an exception helps — it needs to run at that size, and that requires the flattened-offset arithmetic not to overflow. I'm fully on board with your instinct that silent is the worst case: making overflow throw instead of warn is strictly better where we can't support the size (e.g. a device without shaderInt64). But where we can support it, the fix lets the workload actually execute rather than failing loudly.

So the two motivations are complementary:

  1. Concrete capability — a real >32k-env workload that overflows INT32_MAX on main and works with this PR.
  2. Latent memory-safety — the silent wrong-address / under-allocation bug in the shared LLVM path (the point you flagged as a good one), which the same change closes for CPU/CUDA/AMDGPU.

And the SPIR-V commit (55ab538) extends the same narrow widening to Vulkan/Metal (capability-gated on shaderInt64, i32 + warning fallback otherwise), which is the cross-backend parity you asked for — so no backend is left silently overflowing.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

hmmm, this is ndarray only. what happens if we change between ndarray and field? (we do this in genesis-world for example: ndarray for development mode, compiles quickly; field for performance mode, runs quickly, but compiles slowly). Will this introduce a correctness disparity between ndarray and field?

warnings.warn("Ndarray index might be out of int32 boundary but int64 indexing is not supported yet.")
needed_arg_dtype = needed_arg_type.dtype
if needed_arg_dtype is None or id(needed_arg_dtype) in primitive_types.type_ids:
element_dim = 0
else:
element_dim = needed_arg_dtype.ndim
array_shape = v.shape[element_dim:] if is_soa else v.shape[:-element_dim]
if any(dim > np.iinfo(np.int32).max for dim in array_shape):

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Keep product warning for SPIR-V backends

When this runs on Vulkan/Metal, ndarray addressing is still flattened in SPIR-V as an i32 value (quadrants/codegen/spirv/spirv_codegen.cpp:864-894), so a shape like (65536, 65536) has no single dimension above int32 but still wraps the linear offset and reads/writes the wrong element. This branch handles external NumPy/Torch arrays for every arch, so replacing the old product check with only a per-dimension check removes the only warning for unsupported SPIR-V launches; please keep the product warning for non-LLVM backends or widen the SPIR-V path too.

Useful? React with 👍 / 👎.

warnings.warn("Ndarray dimensions above int32 are not supported yet.")
elif (
impl.current_cfg().arch not in _ARCHS_WITH_I64_LINEAR_INDEX
and math.prod(v.shape) > np.iinfo(np.int32).max
and not impl.get_runtime().prog.get_device_caps().get(_qd_core.DeviceCapability.spirv_has_int64)
):
# No single dimension overflows int32, but the flattened element count does. The LLVM backends
# accumulate the linear index in int64, and the SPIR-V backends do too when the device advertises
# shaderInt64. Only a SPIR-V device without shaderInt64 still flattens in int32 and would overflow,
# so warn (this is a cold path -- only huge ndarrays on a non-i64 device reach here).
warnings.warn(
"Ndarray total element count exceeds the int32 boundary; this device lacks 64-bit integer "
"support (shaderInt64), so the linear index is computed in int32 and may overflow. int64 "
"linear indexing requires the LLVM backends (CPU/CUDA/AMDGPU) or a SPIR-V device with shaderInt64."
)
if isinstance(v, np.ndarray):
# Check ndarray flags is expensive (~250ns), so it is important to order branches according to hit stats
if v.flags.c_contiguous:
Expand Down
21 changes: 14 additions & 7 deletions quadrants/codegen/llvm/codegen_llvm.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1773,26 +1773,33 @@ void TaskCodeGenLLVM::visit(ExternalPtrStmt *stmt) {
int num_array_args = num_indices - num_element_indices;
const size_t element_shape_index_offset = num_array_args;

auto *i64_ty = llvm::Type::getInt64Ty(*llvm_context);
for (int i = 0; i < num_array_args; i++) {
auto raw_arg = builder->CreateGEP(
struct_type, llvm_val[stmt->base_ptr],
{tlctx->get_constant(0), tlctx->get_constant(TypeFactory::SHAPE_POS_IN_NDARRAY), tlctx->get_constant(i)});
raw_arg = builder->CreateLoad(tlctx->get_data_type(PrimitiveType::i32), raw_arg);
sizes[i] = raw_arg;
{tlctx->get_constant(0),
tlctx->get_constant(TypeFactory::SHAPE_POS_IN_NDARRAY),
tlctx->get_constant(i)});
raw_arg =
builder->CreateLoad(tlctx->get_data_type(PrimitiveType::i32), raw_arg);
sizes[i] = builder->CreateSExt(raw_arg, i64_ty);
}

auto linear_index = tlctx->get_constant(0);
auto linear_index = tlctx->get_constant(get_data_type<int64>(), 0);
size_t size_var_index = 0;
for (int i = 0; i < num_indices; i++) {
if (i >= element_shape_index_offset && i < element_shape_index_offset + num_element_indices) {
// Indexing TensorType-elements
llvm::Value *size_var = tlctx->get_constant(stmt->element_shape[i - element_shape_index_offset]);
llvm::Value *size_var = tlctx->get_constant(
get_data_type<int64>(),
stmt->element_shape[i - element_shape_index_offset]);
linear_index = builder->CreateMul(linear_index, size_var);
} else {
// Indexing array dimensions
linear_index = builder->CreateMul(linear_index, sizes[size_var_index++]);
}
linear_index = builder->CreateAdd(linear_index, llvm_val[stmt->indices[i]]);
auto index = builder->CreateSExtOrBitCast(llvm_val[stmt->indices[i]], i64_ty);
linear_index = builder->CreateAdd(linear_index, index);
Comment on lines +1801 to +1802

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Document backend-specific large-array indexing

This change expands the user-visible ndarray indexing range beyond INT32_MAX on LLVM backends while leaving Vulkan/Metal limited to int32 linear offsets, but it makes no corresponding change under docs/. The repository's AGENTS.md explicitly requires documentation for public API or usage changes; add end-user documentation describing the newly supported large-array behavior and the backend-specific limitation.

Useful? React with 👍 / 👎.

}
QD_ASSERT(size_var_index == num_indices - num_element_indices);

Expand All @@ -1808,7 +1815,7 @@ void TaskCodeGenLLVM::visit(ExternalPtrStmt *stmt) {
if (operand_dtype->is<TensorType>()) {
// Access PtrOffset via: base_ptr + offset * sizeof(element)

auto address_offset = builder->CreateSExt(linear_index, llvm::Type::getInt64Ty(*llvm_context));
auto address_offset = linear_index;

auto stmt_ret_type = stmt->ret_type.ptr_removed();
if (stmt_ret_type->is<TensorType>()) {
Expand Down
28 changes: 19 additions & 9 deletions quadrants/codegen/spirv/spirv_codegen.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -861,7 +861,13 @@ void TaskCodegen::visit(ExternalTensorShapeAlongAxisStmt *stmt) {
void TaskCodegen::visit(ExternalPtrStmt *stmt) {
// Used mostly for transferring data between host (e.g. numpy array) and
// device.
spirv::Value linear_offset = ir_->int_immediate_number(ir_->i32_type(), 0);
// Flatten the multi-dimensional index into a linear (byte) offset. When the device advertises 64-bit integers we
// accumulate in i64 so that ndarrays whose element count exceeds INT32_MAX do not wrap and silently address the
// wrong element (matches the int64 fix already applied on the LLVM backends). Devices without shaderInt64 keep the
// historical i32 arithmetic; the Python/C++ launch path emits an overflow warning for those.
const bool index_use_i64 = caps_->get(DeviceCapability::spirv_has_int64);
const spirv::SType index_type = index_use_i64 ? ir_->i64_type() : ir_->i32_type();
spirv::Value linear_offset = ir_->int_immediate_number(index_type, 0);
const auto *argload = stmt->base_ptr->as<ArgLoadStmt>();
const auto arg_id = argload->arg_id;
{
Expand All @@ -885,17 +891,18 @@ void TaskCodegen::visit(ExternalPtrStmt *stmt) {
spirv::Value size_var;
// Use immediate numbers to flatten index for element shapes.
if (i >= element_shape_index_offset && i < element_shape_index_offset + element_shape.size()) {
size_var = ir_->uint_immediate_number(ir_->i32_type(), element_shape[i - element_shape_index_offset]);
size_var = ir_->uint_immediate_number(index_type, element_shape[i - element_shape_index_offset]);
} else {
size_var = ir_->query_value(size_var_names[size_var_names_idx++]);
// Shapes are stored as i32 in the args buffer; widen to the accumulation type when using i64.
size_var = ir_->cast(index_type, ir_->query_value(size_var_names[size_var_names_idx++]));
}
spirv::Value indices = ir_->query_value(stmt->indices[i]->raw_name());
spirv::Value indices = ir_->cast(index_type, ir_->query_value(stmt->indices[i]->raw_name()));
linear_offset = ir_->mul(linear_offset, size_var);
linear_offset = ir_->add(linear_offset, indices);
}
size_t type_size = ir_->get_primitive_type_size(stmt->ret_type.ptr_removed());
linear_offset = ir_->make_value(spv::OpShiftLeftLogical, ir_->i32_type(), linear_offset,
ir_->int_immediate_number(ir_->i32_type(), log2int(type_size)));
linear_offset = ir_->make_value(spv::OpShiftLeftLogical, index_type, linear_offset,
ir_->int_immediate_number(index_type, log2int(type_size)));
if (caps_->get(DeviceCapability::spirv_has_no_integer_wrap_decoration)) {
ir_->decorate(spv::OpDecorate, linear_offset, spv::DecorationNoSignedWrap);
}
Expand All @@ -908,7 +915,8 @@ void TaskCodegen::visit(ExternalPtrStmt *stmt) {
spirv::Value addr_ptr = ir_->make_access_chain(ir_->get_pointer_type(ir_->u64_type(), spv::StorageClassUniform),
get_buffer_value(BufferType::Args, PrimitiveType::i32), indices);
spirv::Value base_addr = ir_->load_variable(addr_ptr, ir_->u64_type());
spirv::Value addr = ir_->add(base_addr, ir_->make_value(spv::OpSConvert, ir_->u64_type(), linear_offset));
// cast() sign-extends an i32 offset to u64, or bitcasts an already-64-bit i64 offset to u64.
spirv::Value addr = ir_->add(base_addr, ir_->cast(ir_->u64_type(), linear_offset));
ir_->register_value(stmt->raw_name(), addr);

// Save decomposed base pointer and element index so at_buffer() can
Expand All @@ -917,8 +925,10 @@ void TaskCodegen::visit(ExternalPtrStmt *stmt) {
// per-element reinterpret_cast from ulong arithmetic is miscompiled
// when the stored value is loop-invariant.
size_t type_size = ir_->get_primitive_type_size(stmt->ret_type.ptr_removed());
spirv::Value elem_index = ir_->make_value(spv::OpShiftRightLogical, ir_->i32_type(), linear_offset,
ir_->int_immediate_number(ir_->i32_type(), log2int(type_size)));
// Keep the element index in the same width as the offset accumulation so OpPtrAccessChain indexes the correct
// element for >INT32_MAX-element ndarrays on int64-capable devices.
spirv::Value elem_index = ir_->make_value(spv::OpShiftRightLogical, index_type, linear_offset,
ir_->int_immediate_number(index_type, log2int(type_size)));
physical_ptr_components_[stmt] = {base_addr, elem_index};
} else {
ir_->register_value(stmt->raw_name(), linear_offset);
Expand Down
31 changes: 18 additions & 13 deletions quadrants/program/ndarray.cpp
Original file line number Diff line number Diff line change
@@ -1,8 +1,11 @@
#include <limits>
#include <numeric>

#include "quadrants/common/exceptions.h"
#include "quadrants/program/adstack_size_expr_eval.h"
#include "quadrants/program/ndarray.h"
#include "quadrants/program/program.h"
#include "quadrants/rhi/arch.h"
#include "fp16.h"

#ifdef QD_WITH_LLVM
Expand Down Expand Up @@ -32,7 +35,7 @@ Ndarray::Ndarray(Program *prog,
shape(shape_),
layout(layout_),
dbg_info(dbg_info_),
nelement_(std::accumulate(std::begin(shape_), std::end(shape_), 1, std::multiplies<>())),
nelement_(std::accumulate(std::begin(shape_), std::end(shape_), (std::size_t)1, std::multiplies<>())),
element_size_(data_type_size(dtype)),
prog_(prog) {
// Now that we have two shapes which may be concatenated differently
Expand All @@ -44,11 +47,19 @@ Ndarray::Ndarray(Program *prog,
} else if (layout == ExternalArrayLayout::kSOA) {
total_shape_.insert(total_shape_.begin(), element_shape.begin(), element_shape.end());
}
auto total_num_scalar = std::accumulate(std::begin(total_shape_), std::end(total_shape_), 1LL, std::multiplies<>());
if (total_num_scalar > std::numeric_limits<int>::max()) {
ErrorEmitter(QuadrantsIndexWarning(), &dbg_info,
"Ndarray index might be out of int32 boundary but int64 indexing is "
"not supported yet.");
// The ndarray linear offset in TaskCodegen::visit(ExternalPtrStmt) is flattened in int64 on the LLVM backends
// (CPU/CUDA/AMDGPU) and, on the SPIR-V backends (Vulkan/Metal/...), whenever the device advertises 64-bit
// integers (DeviceCapability::spirv_has_int64). Only a SPIR-V device *without* shaderInt64 still flattens in
// int32, so an owned ndarray whose total element count exceeds int32 would overflow there. Scope the warning to
// exactly that case.
if (!arch_uses_llvm(prog->compile_config().arch) && !prog->get_device_caps().get(DeviceCapability::spirv_has_int64)) {
auto total_num_scalar = std::accumulate(std::begin(total_shape_), std::end(total_shape_), 1LL, std::multiplies<>());
if (total_num_scalar > std::numeric_limits<int>::max()) {
ErrorEmitter(QuadrantsIndexWarning(), &dbg_info,
"Ndarray total element count exceeds the int32 boundary; this device lacks 64-bit integer "
"support (shaderInt64), so the linear index is computed in int32 and may overflow. int64 linear "
"indexing requires the LLVM backends (CPU/CUDA/AMDGPU) or a SPIR-V device with shaderInt64.");
}
}
ndarray_alloc_ = prog->allocate_memory_on_device(nelement_ * element_size_, prog->result_buffer);
}
Expand All @@ -63,7 +74,7 @@ Ndarray::Ndarray(DeviceAllocation &devalloc,
shape(shape),
layout(layout),
dbg_info(dbg_info),
nelement_(std::accumulate(std::begin(shape), std::end(shape), 1, std::multiplies<>())),
nelement_(std::accumulate(std::begin(shape), std::end(shape), (std::size_t)1, std::multiplies<>())),
element_size_(data_type_size(dtype)) {
// When element_shape is specified but layout is not, default layout is AOS.
auto element_shape = data_type_shape(dtype);
Expand All @@ -78,12 +89,6 @@ Ndarray::Ndarray(DeviceAllocation &devalloc,
} else if (layout == ExternalArrayLayout::kSOA) {
total_shape_.insert(total_shape_.begin(), element_shape.begin(), element_shape.end());
}
auto total_num_scalar = std::accumulate(std::begin(total_shape_), std::end(total_shape_), 1LL, std::multiplies<>());
if (total_num_scalar > std::numeric_limits<int>::max()) {
ErrorEmitter(QuadrantsIndexWarning(), &dbg_info,
"Ndarray index might be out of int32 boundary but int64 indexing is "
"not supported yet.");
}
}

Ndarray::Ndarray(DeviceAllocation &devalloc,
Expand Down
90 changes: 90 additions & 0 deletions tests/python/test_ndarray_indexing_i64.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,90 @@
import re

import numpy as np
import psutil
import pytest

import quadrants as qd

from tests import test_utils

# ExternalPtrStmt flattens an N-D ndarray access into a single linear element
# index: for a 2-D array of shape (D0, D1) the codegen emits
# linear = i * D1 + j
# (see TaskCodeGenLLVM::visit(ExternalPtrStmt) in codegen/llvm/codegen_llvm.cpp).
# Before this fix that accumulation was done in i32 and only sign-extended to
# i64 for the final GEP, so any array with more than 2**31 elements overflowed
# *before* the extend and produced a wrong (often negative) address.
#
# The smallest 2-D shape that pushes the last valid linear index past INT32_MAX:
# shape = (2, 2**30 + 1) -> last index [1, 2**30] -> linear = 2**31 + 1
_D1 = 2**30 + 1
_I = 1
_J = 2**30
_TRUE_LINEAR = _I * _D1 + _J # == 2**31 + 1, exceeds np.iinfo(np.int32).max


@test_utils.test(arch=[qd.cpu])
def test_i32_linear_index_overflows_but_i64_is_correct():
# Reproduces the exact arithmetic ExternalPtrStmt performs. The i32 kernel
# mirrors the pre-fix codegen and must wrap to a wrong value; the i64 kernel
# mirrors the fix and must stay correct.
@qd.kernel
def linear_i32(d1: qd.i32, i: qd.i32, j: qd.i32) -> qd.i32:
return i * d1 + j

@qd.kernel
def linear_i64(d1: qd.i64, i: qd.i64, j: qd.i64) -> qd.i64:
return i * d1 + j

assert _TRUE_LINEAR > np.iinfo(np.int32).max

# Pre-fix behavior: i32 accumulation overflows and wraps to a negative offset.
overflowed = linear_i32(_D1, _I, _J)
assert overflowed != _TRUE_LINEAR
assert overflowed < 0

# Post-fix behavior: i64 accumulation yields the true linear index.
assert linear_i64(_D1, _I, _J) == _TRUE_LINEAR


# ~2 GB for the int8 backing array plus headroom for the device-side copy.
_REQUIRED_BYTES = 5 * 1024**3


@pytest.mark.skipif(
psutil.virtual_memory().available < _REQUIRED_BYTES,
reason="needs >5 GB RAM to allocate an ndarray with more than 2**31 elements",
)
@test_utils.test(arch=[qd.cpu])
def test_ndarray_read_past_int32_index_boundary():
# End-to-end regression guard: on the pre-fix codegen the i32 linear index
# for [1, 2**30] wraps negative and this read returns garbage / segfaults.
@qd.kernel
def read(arr: qd.types.NDArray[qd.i8, 2], i: qd.i32, j: qd.i32) -> qd.i8:
return arr[i, j]

np_arr = np.zeros((2, _D1), dtype=np.int8)
sentinel = np.int8(7)
np_arr[_I, _J] = sentinel

assert read(np_arr, _I, _J) == sentinel


@test_utils.test(arch=[qd.cpu])
def test_ndarray_external_ptr_uses_i64_linear_index():
@qd.kernel
def read(arr: qd.types.NDArray[qd.i32, 2], i: qd.i32, j: qd.i32) -> qd.i32:
return arr[i, j]

np_arr = np.arange(16, dtype=np.int32).reshape(4, 4)
assert read(np_arr, 1, 2) == np_arr[1, 2]

compiled = read._primal._last_compiled_kernel_data
assert compiled is not None

llvm_ir = compiled._debug_dump_to_string()

assert re.search(r"sext i32 .* to i64", llvm_ir)
assert "mul nsw i64" in llvm_ir
assert re.search(r"getelementptr i32, ptr .* i64 ", llvm_ir)