66from sparsediffpy ._core ._shapes import validate_shape
77
88
9+ class DimensionError (ValueError ):
10+ """Raised when a value has the wrong number of elements."""
11+ pass
12+
13+
914class Variable (Expression ):
1015 """A decision variable in the expression tree.
1116
@@ -20,61 +25,41 @@ def __init__(self, scope, var_id, shape):
2025
2126 @property
2227 def value (self ):
23- size = self .shape [0 ] * self .shape [1 ]
24- return self ._scope ._flat_values [self ._var_id :self ._var_id + size ].copy ()
28+ return self ._scope ._flat_values [self ._var_id :self ._var_id + self .size ].copy ()
2529
2630 @value .setter
2731 def value (self , val ):
2832 val = np .asarray (val , dtype = np .float64 ).ravel ()
29- size = self .shape [0 ] * self .shape [1 ]
30- if val .size != size :
31- raise ValueError (
32- f"Expected { size } elements for Variable with shape { self .shape } , "
33- f"got { val .size } "
34- )
35- self ._scope ._flat_values [self ._var_id :self ._var_id + size ] = val
33+ if val .size != self .size :
34+ raise DimensionError (f"expected { self .size } elements, got { val .size } " )
35+ self ._scope ._flat_values [self ._var_id :self ._var_id + self .size ] = val
3636
3737
3838class Parameter (Expression ):
3939 """An updatable parameter in the expression tree.
4040
41- Created by Scope.Parameter(). Values are stored on the parameter itself
42- (not in the scope's flat buffer). Updated via .value property .
41+ Created by Scope.Parameter(). Values must be set via .value before
42+ evaluating any expression that uses this parameter .
4343 """
4444
45- def __init__ (self , scope , param_id , shape , value = None ):
45+ def __init__ (self , scope , param_id , shape ):
4646 self ._scope = scope
4747 self ._param_id = param_id
4848 self .shape = shape
49- size = shape [0 ] * shape [1 ]
50- if value is not None :
51- self ._value_flat = np .asarray (value , dtype = np .float64 ).ravel (order = "F" )
52- if self ._value_flat .size != size :
53- raise ValueError (
54- f"Parameter value has { self ._value_flat .size } elements, "
55- f"expected { size } for shape { shape } "
56- )
57- else :
58- self ._value_flat = np .zeros (size , dtype = np .float64 )
49+ self ._value_flat = None
5950
6051 @property
6152 def value (self ):
53+ if self ._value_flat is None :
54+ return None
6255 return self ._value_flat .copy ()
6356
6457 @value .setter
6558 def value (self , val ):
6659 val = np .asarray (val , dtype = np .float64 ).ravel (order = "F" )
67- size = self .shape [0 ] * self .shape [1 ]
68- if val .size != size :
69- raise ValueError (
70- f"Expected { size } elements for Parameter with shape { self .shape } , "
71- f"got { val .size } "
72- )
73- self ._value_flat [:] = val
74-
75-
76- # Patch _is_param_like to recognize Parameter
77- # (already handled via lazy import in _expressions.py)
60+ if val .size != self .size :
61+ raise DimensionError (f"expected { self .size } elements, got { val .size } " )
62+ self ._value_flat = val .copy ()
7863
7964
8065class Scope :
@@ -103,26 +88,27 @@ def Variable(self, d1, d2):
10388 self ._variables .append (var )
10489 return var
10590
106- def Parameter (self , d1 , d2 , value = None ):
107- """Create a new updatable parameter in this scope."""
91+ def Parameter (self , d1 , d2 ):
92+ """Create a new updatable parameter in this scope.
93+
94+ Set its value via .value = ... before evaluating.
95+ """
10896 validate_shape (d1 , d2 )
10997 size = d1 * d2
11098 param_id = self ._next_param_offset
11199 self ._next_param_offset += size
112100
113- param = Parameter (self , param_id , (d1 , d2 ), value )
101+ param = Parameter (self , param_id , (d1 , d2 ))
114102 self ._parameters .append (param )
115103 return param
116104
117- def set_values (self , flat_array ):
105+ def set_values (self , array ):
118106 """Set all variable values at once from a flat array."""
119- flat_array = np .asarray (flat_array , dtype = np .float64 )
120- if flat_array .size != self ._flat_values .size :
121- raise ValueError (
122- f"Expected flat array of size { self ._flat_values .size } , "
123- f"got { flat_array .size } "
124- )
125- self ._flat_values [:] = flat_array
107+ array = np .asarray (array , dtype = np .float64 )
108+ in_size = self ._flat_values .size
109+ if array .size != in_size :
110+ raise DimensionError (f"expected { in_size } elements, got { array .size } " )
111+ self ._flat_values [:] = array
126112
127113 def get_values (self ):
128114 """Return a copy of the flat value buffer."""
0 commit comments