fix: update MCMC and DE tutorials for v1.0.0 API - lambda closures, parameter renames, tuple unpacking
This commit is contained in:
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
+10
-10
@@ -47,13 +47,13 @@ except ImportError:
|
|||||||
|
|
||||||
|
|
||||||
def mcmc_sample(
|
def mcmc_sample(
|
||||||
log_likelihood_fn: Callable[[List[float], List[float]], float],
|
log_likelihood_fn: Callable[[List[float]], float], # Updated: closure captures data
|
||||||
data: np.ndarray,
|
|
||||||
initial_params: np.ndarray,
|
initial_params: np.ndarray,
|
||||||
param_bounds: List[Tuple[float, float]],
|
param_bounds: List[Tuple[float, float]],
|
||||||
n_samples: int = 10000,
|
n_samples: int = 10000,
|
||||||
burn_in: int = 1000,
|
burn_in: int = 1000,
|
||||||
proposal_std: float = 0.1,
|
proposal_std: float = 0.1,
|
||||||
|
data: Optional[np.ndarray] = None, # Deprecated: use closure instead
|
||||||
) -> np.ndarray:
|
) -> np.ndarray:
|
||||||
"""
|
"""
|
||||||
MCMC Metropolis-Hastings sampler.
|
MCMC Metropolis-Hastings sampler.
|
||||||
@@ -65,9 +65,10 @@ def mcmc_sample(
|
|||||||
----------
|
----------
|
||||||
log_likelihood_fn : callable
|
log_likelihood_fn : callable
|
||||||
Function that computes log P(data | params). Should accept
|
Function that computes log P(data | params). Should accept
|
||||||
(params: list, data: list) and return float.
|
(params: list) and return float. Data should be captured in closure.
|
||||||
data : np.ndarray
|
data : np.ndarray, optional (deprecated)
|
||||||
Observed data (passed to log_likelihood_fn)
|
This parameter is deprecated and ignored. Capture data in the
|
||||||
|
log_likelihood_fn closure instead.
|
||||||
initial_params : np.ndarray
|
initial_params : np.ndarray
|
||||||
Starting parameter values
|
Starting parameter values
|
||||||
param_bounds : list of (float, float)
|
param_bounds : list of (float, float)
|
||||||
@@ -103,17 +104,16 @@ def mcmc_sample(
|
|||||||
"""
|
"""
|
||||||
if RUST_AVAILABLE:
|
if RUST_AVAILABLE:
|
||||||
# Convert to lists if numpy arrays
|
# Convert to lists if numpy arrays
|
||||||
data_list = data.tolist() if hasattr(data, 'tolist') else list(data)
|
# Note: data is now captured in log_likelihood_fn closure
|
||||||
params_list = initial_params.tolist() if hasattr(initial_params, 'tolist') else list(initial_params)
|
params_list = initial_params.tolist() if hasattr(initial_params, 'tolist') else list(initial_params)
|
||||||
|
|
||||||
|
# Rust function uses different parameter names
|
||||||
samples = _rust_mcmc_sample(
|
samples = _rust_mcmc_sample(
|
||||||
log_likelihood_fn=log_likelihood_fn,
|
log_likelihood_fn=log_likelihood_fn,
|
||||||
data=data_list,
|
initial_state=params_list, # Rust expects 'initial_state'
|
||||||
initial_params=params_list,
|
|
||||||
param_bounds=param_bounds,
|
|
||||||
n_samples=n_samples,
|
n_samples=n_samples,
|
||||||
|
step_size=proposal_std, # Rust expects 'step_size'
|
||||||
burn_in=burn_in,
|
burn_in=burn_in,
|
||||||
proposal_std=proposal_std,
|
|
||||||
)
|
)
|
||||||
return np.array(samples)
|
return np.array(samples)
|
||||||
else:
|
else:
|
||||||
|
|||||||
Reference in New Issue
Block a user