@@ -201,17 +201,23 @@ def max_aposteriori(self) -> None:
201201 for i , name in enumerate (self .parameter_names ):
202202 self .data_group [name ] = result [i ] * self .parameter_units [i ]
203203
204- def mcmc (self , n_samples : int = 1000 , n_walkers : int = 32 , n_burn : int = 500 , n_thin = 10 ) -> None :
204+ def mcmc (self , x0 : tuple [ sc . Variable ] = None , n_samples : int = 1000 , n_walkers : int = 32 , n_burn : int = 500 , n_thin = 10 ) -> None :
205205 """
206206 Perform MCMC sampling of the model parameters.
207207
208- :param n_samples: Number of samples to generate
209- :param n_walkers: Number of MCMC walkers
210- :param n_burn: Number of burn-in samples
211- :param n_thin: Thinning factor
208+ :param x0: Initial starting position for MCMC sampling. Optional, defaults to max likelihood/aposteriori or mean of samples.
209+ :param n_samples: Number of samples to generate. Optional, defaults to 1000.
210+ :param n_walkers: Number of MCMC walkers. Optional, defaults to 32.
211+ :param n_burn: Number of burn-in samples. Optional, defaults to 500.
212+ :param n_thin: Thinning factor. Optional, defaults to 10.
212213 """
213214 if isinstance (self .data_group [self .parameter_names [0 ]], Samples ):
214215 values = np .array ([sc .mean (self .data_group [p ]).value for p in self .parameter_names ])
216+ elif x0 is not None :
217+ for i , x in enumerate (x0 ):
218+ if self .data_group [self .parameter_names [i ]].unit != x .unit :
219+ raise TypeError ("x0 input must have same unit as paramters" )
220+ values = np .array ([x .value for x in x0 ])
215221 else :
216222 values = np .array ([self .data_group [p ].value for p in self .parameter_names ])
217223 pos = values + values * 1e-2 * np .random .randn (n_walkers , len (self .parameter_names ))
0 commit comments