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
10 changes: 10 additions & 0 deletions src/slangpy_ext/utils/slangpy.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -256,6 +256,16 @@ void NativeBoundVariableRuntime::write_shader_cursor_pre_dispatch(
// We have children, so generate call data for each child and
// store in a dictionary, then store the dictionary as the call data.
ShaderCursor child_field = cursor[m_variable_name.c_str()];
// A reference-typed field is a ConstantBuffer/ParameterBlock sub-object -
// e.g. Slang's CUDA target passes an entry-point uniform struct containing a
// fixed-size descriptor array by reference as an implicit ParameterBlock.
// Dereference before recursing so that children see a cursor whose
// shader_object() owns the offsets they extract: field lookups on a
// reference cursor auto-dereference (yielding sub-object-relative offsets),
// so a child that cached those offsets but wrote through the parent's
// shader object would silently corrupt memory.
if (child_field.is_reference())
child_field = child_field.dereference();
Comment thread
szihs marked this conversation as resolved.
for (const auto& [name, child_ref] : *m_children) {
if (child_ref) {
nb::object child_value = value[name.c_str()];
Expand Down
16 changes: 16 additions & 0 deletions src/slangpy_ext/utils/slangpytensor.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -214,6 +214,22 @@ void TensorMarshall::ensure_binding_info_cached(ShaderCursor cursor, NativeBound
{
if (!m_cached_binding_info.primal.is_valid) {
ShaderCursor field = cursor[binding->variable_name()];
// The cached-offset fast path below assumes the tensor's fields live in
// `cursor.shader_object()` at offsets relative to that object. A
// reference-typed field (a ConstantBuffer/ParameterBlock sub-object) breaks
// that assumption: nested field lookups auto-dereference into the
// sub-object, so the cached offsets would be sub-object-relative while the
// write targets the parent object - silent corruption. Tensor types are
// never passed by reference themselves, and reference-typed *enclosing*
// structs are dereferenced before recursion (see
// NativeBoundVariableRuntime::write_shader_cursor_pre_dispatch), so fail
// loudly if one ever reaches this point.
SGL_CHECK(
!field.is_reference(),
"Tensor binding '{}' is reference-typed (a parameter-group sub-object); "
"the cached-offset writer does not support this shape",
binding->variable_name()
);
m_cached_binding_info = extract_binding_info(field);
}
}
Expand Down
9 changes: 9 additions & 0 deletions src/slangpy_ext/utils/slangpytorchtensor.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -222,6 +222,15 @@ void NativeTorchTensorMarshall::ensure_binding_info_cached(
{
if (!m_cached_binding_info.primal.is_valid) {
ShaderCursor field = cursor[binding->variable_name()];
// See TensorMarshall::ensure_binding_info_cached: the cached-offset writer
// requires the field's offsets to be relative to `cursor.shader_object()`,
// which a reference-typed (parameter-group sub-object) field violates.
SGL_CHECK(
!field.is_reference(),
"Torch tensor binding '{}' is reference-typed (a parameter-group sub-object); "
"the cached-offset writer does not support this shape",
binding->variable_name()
);
m_cached_binding_info = TensorMarshall::extract_binding_info(field);

// Determine copy-back flags from the Slang uniform type name.
Expand Down
25 changes: 22 additions & 3 deletions src/slangpy_ext/utils/slangpyvalue.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -16,8 +16,21 @@ void NativeValueMarshall::ensure_cached(ShaderCursor cursor, NativeBoundVariable
{
if (m_cached.is_valid)
return;
ShaderCursor field
= binding->direct_bind() ? cursor[binding->variable_name()] : cursor[binding->variable_name()]["value"];
ShaderCursor field = cursor[binding->variable_name()];
// A reference-typed field is a ConstantBuffer/ParameterBlock sub-object - e.g.
// Slang's CUDA target passes an entry-point uniform struct carrying a fixed-size
// descriptor array (such as the vectorized-array wrapper Array1DValueType holding
// tensors) by reference. Nested lookups below auto-dereference into the
// sub-object, making every cached offset sub-object-relative, so the write must
// target the sub-object's ShaderObject (see write_shader_cursor_pre_dispatch);
// writing through the parent object with these offsets would corrupt memory.
m_cached.field_is_reference = field.is_reference();
if (m_cached.field_is_reference) {
m_cached.field_index = cursor.find_field_index(binding->variable_name());
field = field.dereference();
}
if (!binding->direct_bind())
field = field["value"];
m_cached.value_offset = field.offset();
m_cached.value_type_layout = field.slang_type_layout();
m_cached.writer = get_shader_cursor_writer(m_cached.value_type_layout);
Expand All @@ -38,7 +51,13 @@ void NativeValueMarshall::write_shader_cursor_pre_dispatch(
AccessType primal_access = binding->access().first;
if (!value.is_none() && (primal_access == AccessType::read || primal_access == AccessType::readwrite)) {
ensure_cached(cursor, binding);
ShaderCursor value_cursor(cursor.shader_object(), m_cached.value_type_layout, m_cached.value_offset);
// For a reference-typed field the cached offsets are relative to the
// sub-object; re-resolve it (cheap: cached field index + object lookup) and
// write there instead of into the parent object.
ShaderObject* target_object = cursor.shader_object();
if (m_cached.field_is_reference)
target_object = cursor.get_field_by_index(m_cached.field_index).dereference().shader_object();
ShaderCursor value_cursor(target_object, m_cached.value_type_layout, m_cached.value_offset);
if (m_cached.writer) {
m_cached.writer(value_cursor, value);
} else {
Expand Down
6 changes: 6 additions & 0 deletions src/slangpy_ext/utils/slangpyvalue.h
Original file line number Diff line number Diff line change
Expand Up @@ -35,6 +35,12 @@ class NativeValueMarshall : public NativeMarshall {
slang::TypeLayoutReflection* value_type_layout = nullptr; ///< Type layout for value field.
std::function<void(ShaderCursor&, nb::object)> writer; ///< Pre-resolved writer fn.
bool direct_bind{false}; ///< direct_bind value used when populating cache.
/// True when the bound field is reference-typed (a ConstantBuffer/ParameterBlock
/// sub-object, e.g. Slang's CUDA by-reference ABI for descriptor-table-carrying
/// uniforms). value_offset is then relative to the sub-object, and writes must
/// target the sub-object's ShaderObject rather than the parent's.
bool field_is_reference{false};
int32_t field_index{-1}; ///< Cached field index for the per-dispatch dereference.
bool is_valid = false;
};

Expand Down
Loading