@@ -1014,71 +1014,122 @@ def test_periodic_scheduler_integration(self):
10141014
10151015
10161016# ============================================================
1017- # Regression: RandomWalk* must return scalar, not numpy array
1018- # https://github.com/scope-lab-vu/ns_gym/issues/XX
1019- # rng.normal(mu, sigma, 1) returns a 1-element array which
1020- # causes ValueError in downstream dynamics (e.g. Acrobot).
1017+ # Parametrized common-property tests for all single_param UpdateFns
10211018# ============================================================
10221019
1023- class TestRandomWalkReturnsScalar :
1024- """All RandomWalk* update functions must return a plain Python scalar,
1025- not a numpy array, so they can be used directly in environment dynamics."""
1026-
1027- def test_random_walk_returns_scalar (self ):
1028- fn = RandomWalk (_always_scheduler (), mu = 0 , sigma = 1 , seed = 42 )
1029- param , _ , _ = fn (10.0 , 0 )
1020+ _single_param_factories = [
1021+ pytest .param (lambda s : NoUpdate (s ), id = "NoUpdate" ),
1022+ pytest .param (lambda s : IncrementUpdate (s , k = 1.0 ), id = "IncrementUpdate" ),
1023+ pytest .param (lambda s : DecrementUpdate (s , k = 1.0 ), id = "DecrementUpdate" ),
1024+ pytest .param (lambda s : DeterministicTrend (s , slope = 1.0 ), id = "DeterministicTrend" ),
1025+ pytest .param (lambda s : RandomWalk (s , mu = 0 , sigma = 1 , seed = 42 ), id = "RandomWalk" ),
1026+ pytest .param (
1027+ lambda s : RandomWalkWithDrift (s , alpha = 1.0 , mu = 0 , sigma = 1 , seed = 42 ),
1028+ id = "RandomWalkWithDrift" ,
1029+ ),
1030+ pytest .param (
1031+ lambda s : RandomWalkWithDriftAndTrend (s , alpha = 1.0 , mu = 0 , sigma = 1 , slope = 0.5 , seed = 42 ),
1032+ id = "RandomWalkWithDriftAndTrend" ,
1033+ ),
1034+ pytest .param (
1035+ lambda s : StepWiseUpdate (s , param_list = [float (i ) for i in range (20 )]),
1036+ id = "StepWiseUpdate" ,
1037+ ),
1038+ pytest .param (lambda s : OscillatingUpdate (s , delta = 1.0 ), id = "OscillatingUpdate" ),
1039+ pytest .param (lambda s : ExponentialDecay (s , decay_rate = 0.5 ), id = "ExponentialDecay" ),
1040+ pytest .param (lambda s : GeometricProgression (s , r = 2.0 ), id = "GeometricProgression" ),
1041+ pytest .param (
1042+ lambda s : OrnsteinUhlenbeck (s , theta = 0.5 , mu = 10.0 , sigma = 0.1 , seed = 42 ),
1043+ id = "OrnsteinUhlenbeck" ,
1044+ ),
1045+ pytest .param (
1046+ lambda s : SigmoidTransition (s , a = 1.0 , b = 10.0 , k = 1.0 , t0 = 50 ),
1047+ id = "SigmoidTransition" ,
1048+ ),
1049+ pytest .param (
1050+ lambda s : CyclicUpdate (s , value_list = [10.0 , 20.0 , 30.0 ]),
1051+ id = "CyclicUpdate" ,
1052+ ),
1053+ pytest .param (
1054+ lambda s : BoundedRandomWalk (s , mu = 0 , sigma = 1 , lo = - 100 , hi = 100 , seed = 42 ),
1055+ id = "BoundedRandomWalk" ,
1056+ ),
1057+ pytest .param (lambda s : PolynomialTrend (s , coeffs = [1.0 ]), id = "PolynomialTrend" ),
1058+ pytest .param (
1059+ lambda s : LinearInterpolation (s , start_val = 0.0 , end_val = 10.0 , T = 100 ),
1060+ id = "LinearInterpolation" ,
1061+ ),
1062+ ]
1063+
1064+
1065+ class TestSingleParamCommonProperties :
1066+ """Parametrized tests for properties every single_param UpdateFn must satisfy."""
1067+
1068+ @pytest .mark .parametrize ("make_fn" , _single_param_factories )
1069+ def test_returns_scalar (self , make_fn ):
1070+ """Output param must be a scalar, not a numpy array."""
1071+ fn = make_fn (_always_scheduler ())
1072+ param , _ , _ = fn (10.0 , 1 )
10301073 assert np .isscalar (param ), (
1031- f"RandomWalk returned { type (param )} , expected scalar"
1074+ f"{ fn . __class__ . __name__ } returned { type (param )} , expected scalar"
10321075 )
10331076
1034- def test_random_walk_with_drift_returns_scalar ( self ):
1035- fn = RandomWalkWithDrift (
1036- _always_scheduler (), alpha = 1.0 , mu = 0 , sigma = 1 , seed = 42
1037- )
1038- param , _ , _ = fn (10.0 , 0 )
1039- assert np . isscalar ( param ) , (
1040- f"RandomWalkWithDrift returned { type ( param ) } , expected scalar "
1077+ @ pytest . mark . parametrize ( "make_fn" , _single_param_factories )
1078+ def test_returns_three_tuple ( self , make_fn ):
1079+ """__call__ must return a (param, changed, delta_change) 3-tuple."""
1080+ fn = make_fn ( _always_scheduler () )
1081+ result = fn (10.0 , 1 )
1082+ assert len ( result ) == 3 , (
1083+ f"{ fn . __class__ . __name__ } returned { len ( result ) } -tuple , expected 3 "
10411084 )
10421085
1043- def test_random_walk_with_drift_and_trend_returns_scalar ( self ):
1044- fn = RandomWalkWithDriftAndTrend (
1045- _always_scheduler (), alpha = 1.0 , mu = 0 , sigma = 1 , slope = 0.5 , seed = 42
1046- )
1047- param , _ , _ = fn (10.0 , 1 )
1048- assert np . isscalar ( param ), (
1049- f"RandomWalkWithDriftAndTrend returned { type ( param ) } , expected scalar "
1086+ @ pytest . mark . parametrize ( "make_fn" , _single_param_factories )
1087+ def test_changed_flag_is_zero_or_one ( self , make_fn ):
1088+ """The changed flag must be 0 or 1."""
1089+ fn = make_fn ( _always_scheduler () )
1090+ _ , changed , _ = fn (10.0 , 1 )
1091+ assert changed in ( 0 , 1 ), (
1092+ f"{ fn . __class__ . __name__ } returned changed= { changed } , expected 0 or 1 "
10501093 )
10511094
1052- def test_random_walk_scalar_survives_multiple_steps (self ):
1053- """Ensure the scalar type is preserved across multiple update steps."""
1054- fn = RandomWalk (_always_scheduler (), mu = 0 , sigma = 1 , seed = 42 )
1055- param = 10.0
1056- for t in range (10 ):
1057- param , _ , _ = fn (param , t )
1058- assert np .isscalar (param ), (
1059- f"RandomWalk returned { type (param )} at step { t } , expected scalar"
1060- )
1095+ @pytest .mark .parametrize ("make_fn" , _single_param_factories )
1096+ def test_delta_change_is_numeric (self , make_fn ):
1097+ """The delta_change element must be a number."""
1098+ fn = make_fn (_always_scheduler ())
1099+ _ , _ , delta = fn (10.0 , 1 )
1100+ assert isinstance (delta , (int , float , np .integer , np .floating )), (
1101+ f"{ fn .__class__ .__name__ } returned delta type { type (delta )} , expected numeric"
1102+ )
10611103
1062- def test_random_walk_with_drift_scalar_survives_multiple_steps (self ):
1063- fn = RandomWalkWithDrift (
1064- _always_scheduler (), alpha = 0.5 , mu = 0 , sigma = 1 , seed = 42
1104+ @pytest .mark .parametrize ("make_fn" , _single_param_factories )
1105+ def test_changed_is_one_when_scheduler_fires (self , make_fn ):
1106+ """When the scheduler fires, changed must be 1."""
1107+ fn = make_fn (_always_scheduler ())
1108+ _ , changed , _ = fn (10.0 , 1 )
1109+ assert changed == 1 , (
1110+ f"{ fn .__class__ .__name__ } returned changed={ changed } , expected 1"
10651111 )
1066- param = 10.0
1067- for t in range (10 ):
1068- param , _ , _ = fn (param , t )
1069- assert np .isscalar (param ), (
1070- f"RandomWalkWithDrift returned { type (param )} at step { t } , expected scalar"
1071- )
10721112
1073- def test_random_walk_with_drift_and_trend_scalar_survives_multiple_steps (self ):
1074- fn = RandomWalkWithDriftAndTrend (
1075- _always_scheduler (), alpha = 0.5 , mu = 0 , sigma = 1 , slope = 0.1 , seed = 42
1113+ @pytest .mark .parametrize ("make_fn" , _single_param_factories )
1114+ def test_no_update_when_scheduler_false (self , make_fn ):
1115+ """When the scheduler does not fire: param unchanged, changed=0, delta=0."""
1116+ fn = make_fn (_never_scheduler ())
1117+ param , changed , delta = fn (10.0 , 0 )
1118+ assert param == 10.0 , (
1119+ f"{ fn .__class__ .__name__ } changed param to { param } when scheduler was off"
10761120 )
1121+ assert changed == 0
1122+ assert delta == 0.0
1123+
1124+ @pytest .mark .parametrize ("make_fn" , _single_param_factories )
1125+ def test_scalar_survives_multiple_steps (self , make_fn ):
1126+ """Param must stay scalar across 10 consecutive update steps."""
1127+ fn = make_fn (_always_scheduler ())
10771128 param = 10.0
10781129 for t in range (10 ):
10791130 param , _ , _ = fn (param , t )
10801131 assert np .isscalar (param ), (
1081- f"RandomWalkWithDriftAndTrend returned { type (param )} at step { t } , expected scalar"
1132+ f"{ fn . __class__ . __name__ } returned { type (param )} at step { t } , expected scalar"
10821133 )
10831134
10841135
0 commit comments