diff --git a/conftest.py b/conftest.py index 1697964..6990c1c 100644 --- a/conftest.py +++ b/conftest.py @@ -17,6 +17,7 @@ def pytest_addoption(parser): additional_modules += list((package_loc / "csdl_alpha" / "src" / "operations" / "sparse").glob("*.py")) additional_modules += list((package_loc / "csdl_alpha" / "src" / "operations" / "derivatives").glob("*.py")) additional_modules += list((package_loc / "csdl_alpha" / "src" / "operations" / "special").glob("*.py")) +additional_modules += list((package_loc / "csdl_alpha" / "src" / "operations" / "random").glob("*.py")) def pytest_collect_file(file_path, path, parent): if file_path in additional_modules: diff --git a/csdl_alpha/backends/jax/graph_to_jax.py b/csdl_alpha/backends/jax/graph_to_jax.py index 2917a6e..42d7d62 100644 --- a/csdl_alpha/backends/jax/graph_to_jax.py +++ b/csdl_alpha/backends/jax/graph_to_jax.py @@ -1,14 +1,15 @@ +from multiprocessing import Array from ...src.graph.graph import Graph from ...src.graph.variable import Variable from ...utils.inputs import listify_variables from csdl_alpha.src.operations.loops.loop import Loop from csdl_alpha.src.operations.implicit_operations.implicit_operation import ImplicitOperation +from csdl_alpha.src.operations.operation_subclasses import RandomOperation import numpy as np from typing import Union, Callable - # Get the graph def get_jax_inputs(node, all_jax_variables:dict, return_dict = False)->list: import jax.numpy as jnp @@ -72,6 +73,9 @@ def create_jax_function( """ current_graph = graph import jax.numpy as jnp + from jax import random + from jax._src.typing import Array + KeyArray = Array # Type alias for JAX PRNGKey inputs = list(inputs) outputs = list(outputs) @@ -99,7 +103,22 @@ def create_jax_function( # all_sorted_nodes = build_derivative_node_order(current_graph, outputs, inputs, reverse=False) # Build the JAX function itself - def jax_function(*args)->list: + def jax_function(*args, prng_key:KeyArray=None)->list: + """JAX function that computes the outputs given the inputs. + + Parameters + ---------- + args : list + The input values to the function, in the same order as the inputs. + prng_key : KeyArray, optional + JAX PRNG key for random operations. If the graph contains random operations, this must be provided. + + Returns + ------- + list + The computed outputs, in the same order as the outputs. + """ + # Set the input values all_jax_variables = {} relevant_nodes = set() @@ -140,6 +159,21 @@ def jax_function(*args)->list: if fill_outputs[output_node] is None: raise ValueError(f"Jax function error with node {node.info()}: Output {output_node.info()} was not filled") all_jax_variables[output_node] = fill_outputs[output_node] + elif isinstance(node, RandomOperation): + # Random operation get a key as an argument to their compute_jax function + if prng_key is None: + raise ValueError(f"Jax function error with node {node.info()}: Random operation requires a PRNG key, but got None") + # split the key for each random operation + prng_key, node_key = random.split(prng_key, 2) + + jax_inputs = get_jax_inputs(node, all_jax_variables) + # key is always the first argument for random operations, so we prepend it to jax_inputs + jax_inputs = [node_key] + jax_inputs + jax_outputs = node.compute_jax(*jax_inputs) + if isinstance(jax_outputs, jnp.ndarray): + jax_outputs = (jax_outputs,) + update_jax_variables(node, jax_outputs, all_jax_variables) + else: jax_inputs = get_jax_inputs(node, all_jax_variables) jax_outputs = node.compute_jax(*jax_inputs) # EVERY CSDL OPERATIONS NEEDS THIS FUNCTION @@ -161,7 +195,7 @@ def create_jax_interface( graph:Graph = None, device:str='gpu', enable_f64:bool=True, - name = 'jax_interface')->Callable[[dict[Variable, np.ndarray]], dict[Variable, np.ndarray]]: + name = 'jax_interface'): """_summary_ Parameters @@ -176,10 +210,21 @@ def create_jax_interface( Returns ------- jax interface: Callable - A function with type signature: jax_interface(dict[Variable, np.array])->dict[Variable, np.array], where the input and output variables must match the inputs and outputs respectively. + A function with type signature: + jax_interface(dict[Variable, np.array], prng_key:KeyArray=None) -> dict[Variable, np.array], + where the input and output variables must match the inputs and outputs respectively. + + Notes + ----- + The returned jax_interface function accepts an optional keyword argument: + + prng_key : KeyArray, optional + JAX PRNG key for random operations. Required if the graph contains random operations. """ import jax import csdl_alpha as csdl + from jax._src.typing import Array + KeyArray = Array # Type alias for JAX PRNGKey # import os # os.environ['XLA_FLAGS'] = ( @@ -216,16 +261,29 @@ def create_jax_interface( jax_function = jax.jit(jax_function, device=device) # Create the JAX interface - def jax_interface(inputs_dict:dict[Variable, np.ndarray])->dict[Variable, np.ndarray]: + def jax_interface(inputs_dict:dict[Variable, np.ndarray], prng_key:KeyArray=None)->dict[Variable, np.ndarray]: + """ + Parameters + ---------- + inputs_dict : dict[Variable, np.ndarray] + Dictionary mapping input Variables to their values. + prng_key : KeyArray, optional + JAX PRNG key for random operations. Required if the graph contains random operations. + + Returns + ------- + outputs_dict : dict[Variable, np.ndarray] + Dictionary mapping output Variables to their computed values. + """ jax_interface_inputs = [] # print('INPUTS:') for input_var in inputs: jax_interface_inputs.append(jax.numpy.array(inputs_dict[input_var])) - jax_outputs = jax_function(*jax_interface_inputs) + jax_outputs = jax_function(*jax_interface_inputs, prng_key=prng_key) #### Potential analysis tools #### ## ----- compiled func cost estimates ----- - traced = jax_function.trace(*jax_interface_inputs) + traced = jax_function.trace(*jax_interface_inputs, prng_key=prng_key) lowered = traced.lower() compiled = lowered.compile() # for compiled_costs in compiled.cost_analysis(): diff --git a/csdl_alpha/backends/jax/utils.py b/csdl_alpha/backends/jax/utils.py index 12e3ef1..eceb8a5 100644 --- a/csdl_alpha/backends/jax/utils.py +++ b/csdl_alpha/backends/jax/utils.py @@ -15,7 +15,8 @@ def new_inline_func(*args_in): output = jax.pure_callback( new_inline_func, [jax.ShapeDtypeStruct(output.shape, np.float64) for output in operation.outputs], - *args) + *args, + vmap_method="sequential") return tuple(output) diff --git a/csdl_alpha/src/operations/__init__.py b/csdl_alpha/src/operations/__init__.py index 43e00c7..6684df2 100644 --- a/csdl_alpha/src/operations/__init__.py +++ b/csdl_alpha/src/operations/__init__.py @@ -58,5 +58,9 @@ from .special.bessel import bessel from .special.activations import sigmoid, softplus, relu +# Random operations +from .random.bernoulli import bernoulli +from .random.normal import normal + # other from .subop import subop diff --git a/csdl_alpha/src/operations/custom/custom.py b/csdl_alpha/src/operations/custom/custom.py index 63e7d03..ccd579a 100644 --- a/csdl_alpha/src/operations/custom/custom.py +++ b/csdl_alpha/src/operations/custom/custom.py @@ -97,7 +97,8 @@ def new_inline_func(*args): output = jax.pure_callback( new_inline_func, [jax.ShapeDtypeStruct(self.output_dict[output_var].shape, dtype) for output_var in self.output_dict], - *args) + *args, + vmap_method="sequential") # if len(output) == 1: # output = output[0] return tuple(output) @@ -463,7 +464,8 @@ def new_inline_func(*args): output = jax.pure_callback( new_inline_func, [jax.ShapeDtypeStruct(in_cot.shape, dtype) for in_cot in self.input_cotangents], - *args) + *args, + vmap_method="sequential") # if len(output) == 1: # output = output[0] diff --git a/csdl_alpha/src/operations/operation_subclasses.py b/csdl_alpha/src/operations/operation_subclasses.py index 34d71f9..50077b4 100644 --- a/csdl_alpha/src/operations/operation_subclasses.py +++ b/csdl_alpha/src/operations/operation_subclasses.py @@ -299,6 +299,14 @@ def evaluate_composed(self, *args): inverse = InvertedComposedOperation(y_value, *self.inputs).finalize_and_return_outputs() return inverse +@set_properties() +class RandomOperation(Operation): + """ + Base class for random operations. + """ + def evaluate_vjp(self, *args): + raise NotImplementedError("Random operations do not support VJP evaluation.") + class SubgraphFunctionOperation(SubgraphOperation): def __init__( self, diff --git a/csdl_alpha/src/operations/random/__init__.py b/csdl_alpha/src/operations/random/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/csdl_alpha/src/operations/random/bernoulli.py b/csdl_alpha/src/operations/random/bernoulli.py new file mode 100644 index 0000000..bfc44da --- /dev/null +++ b/csdl_alpha/src/operations/random/bernoulli.py @@ -0,0 +1,115 @@ +import pytest +from csdl_alpha.src.operations.operation_subclasses import RandomOperation +from csdl_alpha.src.graph.variable import Variable +from csdl_alpha.utils.inputs import validate_and_variablize +from csdl_alpha.utils.typing import VariableLike +import csdl_alpha.utils.testing_utils as csdl_tests +import numpy as np + + +class Bernoulli(RandomOperation): + def __init__(self, p:Variable, shape:tuple): + """Initialize the Bernoulli distribution. + + Parameters + ---------- + p : Variable + The probability of success (i.e., the probability of getting 1). + shape : tuple + The shape of the output array. + """ + super().__init__(p) + self.shape = shape + self.name = 'bernoulli' + self.set_dense_outputs((self.shape,)) + + def compute_inline(self, p): + return np.random.binomial(1, p, self.shape) + + def compute_jax(self, key, p): + """Compute the Bernoulli distribution using JAX. + + Parameters + ---------- + key : jax.random.PRNGKey + The random key for JAX. + + Returns + ------- + jax.numpy.ndarray + An array of shape `self.shape` with values drawn from a Bernoulli distribution. + """ + from jax import random + + return random.bernoulli(key, p, shape=self.shape) + + +def bernoulli(p:VariableLike, shape:tuple) -> Variable: + """Create a Bernoulli distribution variable. + + Parameters + ---------- + p : VariableLike + The probability of success (i.e., the probability of getting 1). + shape : tuple + The shape of the output array. + + Returns + ------- + Variable + A variable representing the Bernoulli distribution, with shape `shape`. + """ + p = validate_and_variablize(p) + return Bernoulli(p, shape).finalize_and_return_outputs() + + + +class TestBernoulli(csdl_tests.CSDLTest): + def test_functionality(self): + self.prep(always_build_inline=True) + + import csdl_alpha as csdl + import numpy as np + + p_val = 0.7 + shape = (3, 4) + + p = csdl.Variable(value=p_val) + bernoulli_var = bernoulli(p, shape) + + # Check the shape of the result + assert bernoulli_var.shape == shape, f"Expected shape {shape}, got {bernoulli_var.shape}" + + # Check that the values are either 0 or 1 + assert np.all(np.isin(bernoulli_var.value, [0, 1])), f"Expected values to be 0 or 1, got {bernoulli_var.value}" + + def test_jax_interface(self): + self.prep() + + import csdl_alpha as csdl + import jax.numpy as jnp + from jax import random + + p_val = 0.7 + shape = (3, 4) + + p = csdl.Variable(value=p_val) + bernoulli_var = bernoulli(p, shape) + + interface = csdl.jax.create_jax_interface(p, bernoulli_var) + outputs = interface({p:p.value}, prng_key=random.PRNGKey(42)) + + def test_derivative_error(self): + self.prep() + + import csdl_alpha as csdl + + p_val = 0.7 + shape = (3, 4) + + p = csdl.Variable(value=p_val) + bernoulli_var = bernoulli(p, shape) + + # Check that the derivative raises an error + with pytest.raises(ValueError): + csdl.derivative(bernoulli_var, p) \ No newline at end of file diff --git a/csdl_alpha/src/operations/random/normal.py b/csdl_alpha/src/operations/random/normal.py new file mode 100644 index 0000000..f063f67 --- /dev/null +++ b/csdl_alpha/src/operations/random/normal.py @@ -0,0 +1,65 @@ +import pytest +from csdl_alpha.src.operations.operation_subclasses import RandomOperation +from csdl_alpha.src.graph.variable import Variable +from csdl_alpha.utils.inputs import validate_and_variablize +from csdl_alpha.utils.typing import VariableLike +import csdl_alpha.utils.testing_utils as csdl_tests +import numpy as np + +class Normal(RandomOperation): + def __init__(self, shape:tuple): + super().__init__() + self.shape = shape + self.name = 'normal' + self.set_dense_outputs((self.shape,)) + + def compute_inline(self): + return np.random.normal(size=self.shape) + + def compute_jax(self, key): + """Compute the normal distribution using JAX. + + Parameters + ---------- + key : jax.random.PRNGKey + The random key for JAX. + + Returns + ------- + jax.numpy.ndarray + An array of shape `self.shape` with values drawn from a normal distribution. + """ + from jax import random + return random.normal(key, shape=self.shape) + +def normal(shape:tuple) -> Variable: + """Create a normal distribution variable. + + Parameters + ---------- + shape : tuple + The shape of the output array. + + Returns + ------- + Variable + A variable representing the normal distribution, with shape `shape`. + """ + return Normal(shape).finalize_and_return_outputs() + + +class TestNormal(csdl_tests.CSDLTest): + def test_functionality(self): + self.prep(always_build_inline=True) + + import csdl_alpha as csdl + import numpy as np + + shape = (3, 4) + normal_var = normal(shape) + + # Check the shape of the result + assert normal_var.shape == shape, f"Expected shape {shape}, got {normal_var.shape}" + + # Check that the values are normally distributed + assert np.all(np.isfinite(normal_var.value)), "Expected all values to be finite" diff --git a/csdl_alpha/src/operations/set_get/setindex.py b/csdl_alpha/src/operations/set_get/setindex.py index f375fd1..242d330 100644 --- a/csdl_alpha/src/operations/set_get/setindex.py +++ b/csdl_alpha/src/operations/set_get/setindex.py @@ -31,7 +31,15 @@ def __init__( def compute_inline(self, x, y, *slice_args): x_updated = x.copy() - x_updated[self.slice.evaluate(*slice_args)] = y + eval_slice = self.slice.evaluate(*slice_args) + tgt_shape = np.shape(x_updated[eval_slice]) + # NumPy 2 no longer silently squeezes a size-1 RHS into a size-1 (e.g. scalar) + # target; reshape y to match so the assignment succeeds. Gated on the target + # also being size-1 so a genuinely mismatched y (a real bug elsewhere) still + # raises numpy's broadcast error instead of being silently reshaped. + if np.size(y) == 1 and int(np.prod(tgt_shape)) == 1 and np.shape(y) != tgt_shape: + y = np.asarray(y).reshape(tgt_shape) + x_updated[eval_slice] = y return x_updated # # Set item could add over duplicate indices. diff --git a/csdl_alpha/src/operations/subop.py b/csdl_alpha/src/operations/subop.py index cf923ef..b01a614 100644 --- a/csdl_alpha/src/operations/subop.py +++ b/csdl_alpha/src/operations/subop.py @@ -41,7 +41,8 @@ def compute_jax(self, *args): if self.jit: output = jax.pure_callback(jax.jit(jax_function), [jax.ShapeDtypeStruct(output.shape, np.float64) for output in self.outputs], - *args) + *args, + vmap_method="sequential") return tuple(output) # return tuple(jax.jit(jax_function)(*args)) else: