Skip to content
Merged
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
3 changes: 3 additions & 0 deletions numba_cfunc_compiler/defaults/dict_support.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,6 +37,9 @@ def get_methods(self):
def _to_voidptr_func_name(self) -> str:
return "standalone_dict_to_voidptr"

def _free_func_name(self) -> str:
return "standalone_dict_free"

def _key_type_name(self) -> str:
return NumbaTypeRegistry.resolve_numba_name(self.value.key_type)

Expand Down
3 changes: 3 additions & 0 deletions numba_cfunc_compiler/defaults/list_support.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,9 @@ def get_methods(self):
def _to_voidptr_func_name(self) -> str:
return "standalone_list_to_voidptr"

def _free_func_name(self) -> str:
return "standalone_list_free"

def _elem_type_name(self) -> str:
return NumbaTypeRegistry.resolve_numba_name(self.value.element_type)

Expand Down
14 changes: 13 additions & 1 deletion numba_cfunc_compiler/defaults/primitive_support.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,7 +81,19 @@ def try_parse_state(cls, node: ast.AnnAssign, var_name: str, globalns: dict) ->
if not isinstance(node.value, ast.Constant):
raise TypeError(f"State '{var_name}' must have a literal initial value")

return StateVariableInfo(var_name, node.value.value, state_type)
initial_value = node.value.value
initial_type = type(initial_value)

# State storage is allocated by the host from the concrete Python value,
# while generated code reads it using the declared State type. Keep those
# representations identical, allowing only the safe numeric widening that
# is already supported for inputs.
if state_type is float and initial_type is int:
initial_value = float(initial_value)
elif initial_type is not state_type:
raise TypeError(f"State '{var_name}' expected an initial value of type {state_type.__name__}, got {initial_type.__name__}")

return StateVariableInfo(var_name, initial_value, state_type)


def register():
Expand Down
24 changes: 24 additions & 0 deletions numba_cfunc_compiler/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -269,6 +269,9 @@ def init_statements(self, var_name: str, loaded_value_target, state_slot_target=
def _to_voidptr_func_name(self) -> str:
raise NotImplementedError

def _free_func_name(self) -> str:
raise NotImplementedError

def create_new_container(self, var_name: str) -> list[ast.AST]:
raise NotImplementedError

Expand Down Expand Up @@ -310,6 +313,27 @@ def emit_container_state_load(standalone_state_vars, state_array_name: str = STA
)
return statements

@staticmethod
def emit_container_state_free(standalone_state_vars, state_array_name: str = STATE_ARRAY_NAME) -> list[ast.stmt]:
"""Free state containers and clear their host state slots on STOP."""
statements: list[ast.stmt] = []
for v in standalone_state_vars:
statements.append(
ast.Expr(
value=AST.function_call(
v.type._free_func_name(),
ast.Name(id=v.local_variable_name(), ctx=ast.Load()),
)
)
)
statements.append(
AST.assignment(
AST.array_access(state_array_name, v.array_idx),
AST.function_call("voidptr_null"),
)
)
return statements


@dataclass(frozen=True)
class UnknownType(VariableType):
Expand Down
8 changes: 7 additions & 1 deletion numba_cfunc_compiler/numba_ast_converter.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import ast
import copy

from numba_cfunc_compiler.defaults.struct_support import StructType
from numba_cfunc_compiler.models import (
Expand Down Expand Up @@ -118,7 +119,12 @@ def visit_FunctionDef(self, node):

# Prepend container loading to execution_body
container_load = ContainerType.emit_container_state_load(container_state_vars)
execution_body = container_load + execution_body
execution_body = copy.deepcopy(container_load) + execution_body

# STOP needs typed container values for user cleanup code, then must
# release the native allocations before the host drops its slots.
container_free = ContainerType.emit_container_state_free(container_state_vars)
transformed_stop_body = container_load + transformed_stop_body + container_free

# Build the lifecycle-aware body
lifecycle_body = []
Expand Down
4 changes: 4 additions & 0 deletions numba_cfunc_compiler/numba_core.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,12 +22,14 @@
_standalone_dict_iter_begin,
_standalone_dict_iter_next_item,
_standalone_dict_iter_next_key,
standalone_dict_free,
standalone_dict_from_voidptr,
standalone_dict_length,
standalone_dict_new,
standalone_dict_to_voidptr,
)
from numba_cfunc_compiler.standalone.list import (
standalone_list_free,
standalone_list_from_voidptr,
standalone_list_new,
standalone_list_to_voidptr,
Expand Down Expand Up @@ -266,10 +268,12 @@ def create_compiled_func(
# standalone list (NRT-free)
"standalone_list_new": standalone_list_new,
"standalone_list_from_voidptr": standalone_list_from_voidptr,
"standalone_list_free": standalone_list_free,
"standalone_list_to_voidptr": standalone_list_to_voidptr,
# standalone dict (NRT-free)
"standalone_dict_new": standalone_dict_new,
"standalone_dict_from_voidptr": standalone_dict_from_voidptr,
"standalone_dict_free": standalone_dict_free,
"standalone_dict_to_voidptr": standalone_dict_to_voidptr,
"standalone_dict_length": standalone_dict_length,
"_standalone_dict_iter_begin": _standalone_dict_iter_begin,
Expand Down
4 changes: 4 additions & 0 deletions numba_cfunc_compiler/standalone/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,13 +2,15 @@

from numba_cfunc_compiler.standalone.dict import (
StandaloneDictType,
standalone_dict_free,
standalone_dict_from_voidptr,
standalone_dict_length,
standalone_dict_new,
standalone_dict_to_voidptr,
)
from numba_cfunc_compiler.standalone.list import (
StandaloneListType,
standalone_list_free,
standalone_list_from_voidptr,
standalone_list_length,
standalone_list_new,
Expand All @@ -20,10 +22,12 @@
"StandaloneDictType",
# List
"StandaloneListType",
"standalone_dict_free",
"standalone_dict_from_voidptr",
"standalone_dict_length",
"standalone_dict_new",
"standalone_dict_to_voidptr",
"standalone_list_free",
"standalone_list_from_voidptr",
"standalone_list_length",
"standalone_list_new",
Expand Down
16 changes: 16 additions & 0 deletions numba_cfunc_compiler/standalone/dict.py
Original file line number Diff line number Diff line change
Expand Up @@ -153,6 +153,22 @@ def codegen(context, builder, signature, args):
return sig, codegen


@intrinsic
def standalone_dict_free(typingctx, dict_ty):
"""Release an NB_Dict allocated by :func:`standalone_dict_new`."""
if isinstance(dict_ty, StandaloneDictType):
sig = types.void(dict_ty)

def codegen(context, builder, signature, args):
[dict_ptr] = args
fnty = ir.FunctionType(ir.VoidType(), [i8ptr()])
fn = get_or_declare_function(builder.module, "numba_dict_free", fnty)
builder.call(fn, [dict_ptr])
return context.get_dummy_value()

return sig, codegen


@intrinsic
def standalone_dict_length(typingctx, dict_ty):
"""
Expand Down
16 changes: 16 additions & 0 deletions numba_cfunc_compiler/standalone/list.py
Original file line number Diff line number Diff line change
Expand Up @@ -186,6 +186,22 @@ def codegen(context, builder, signature, args):
return sig, codegen


@intrinsic
def standalone_list_free(typingctx, lst_ty):
"""Release an NB_List allocated by :func:`standalone_list_new`."""
if isinstance(lst_ty, StandaloneListType):
sig = types.void(lst_ty)

def codegen(context, builder, signature, args):
[lst_ptr] = args
fnty = ir.FunctionType(ir.VoidType(), [i8ptr()])
fn = get_or_declare_function(builder.module, "numba_list_free", fnty)
builder.call(fn, [lst_ptr])
return context.get_dummy_value()

return sig, codegen


@intrinsic
def standalone_list_length(typingctx, lst_ty):
"""
Expand Down
12 changes: 12 additions & 0 deletions numba_cfunc_compiler/tests/test_containers.py
Original file line number Diff line number Diff line change
Expand Up @@ -105,6 +105,12 @@ def test_float_values_accumulate(self):
got = [node.execute([k, v])[0] for k, v in zip(keys, vals)]
self.assertEqual(got, [10.0, 20.0, 40.0, 60.0, 90.0])

def test_stop_frees_state(self):
node = CompiledNode(compile_function(dict_contains), input_types=[int]).start()
node.execute([1])
node.stop()
self.assertIsNone(node._state[0])


class TestListExecution(unittest.TestCase):
def test_append_len(self):
Expand All @@ -118,6 +124,12 @@ def test_append_getitem(self):
self.assertEqual(node.execute([10])[0], 10)
self.assertEqual(node.execute([20])[0], 20)

def test_stop_frees_state(self):
node = CompiledNode(compile_function(list_append_len), input_types=[int]).start()
node.execute([10])
node.stop()
self.assertIsNone(node._state[0])


if __name__ == "__main__":
unittest.main()
5 changes: 5 additions & 0 deletions numba_cfunc_compiler/tests/test_support_units.py
Original file line number Diff line number Diff line change
Expand Up @@ -475,6 +475,11 @@ def test_models_type_factory_registry_and_source_registry():
assert TypeFactory.try_parse_input(param, int)[1] == ParameterInfo(int)
assert TypeFactory.try_parse_input(param, str) is None
assert TypeFactory.try_parse_state(parse_stmt("x: State[int] = 1"), "x", {}) == StateVariableInfo("x", 1, int)
assert TypeFactory.try_parse_state(parse_stmt("x: State[float] = 1"), "x", {}) == StateVariableInfo("x", 1.0, float)
assert TypeFactory.try_parse_state(parse_stmt("x: State[bool] = True"), "x", {}) == StateVariableInfo("x", True, bool)
for annotation, value in (("int", "1.0"), ("int", "True"), ("float", "True"), ("bool", "1")):
with pytest.raises(TypeError, match=f"expected an initial value of type {annotation}"):
TypeFactory.try_parse_state(parse_stmt(f"x: State[{annotation}] = {value}"), "x", {})
assert TypeFactory.try_parse_state(parse_stmt("x: State[str] = 'a'"), "x", {}) is None

int_info = NumbaTypeRegistry.get_by_python_type(int)
Expand Down
Loading