Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
64 changes: 45 additions & 19 deletions cmdstanpy/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,6 @@
CmdStanArgs,
GenerateQuantitiesArgs,
LaplaceArgs,
Method,
OptimizeArgs,
PathfinderArgs,
SamplerArgs,
Expand Down Expand Up @@ -434,16 +433,21 @@ def optimize(
)
runset.raise_for_timeouts()

if not runset._check_retcodes():
converged = runset._check_retcodes()
if not converged:
msg = "Error during optimization! Command '{}' failed: {}".format(
' '.join(runset.cmd(0)), runset.get_err_msgs()
)
if 'Line search failed' in msg and not require_converged:
get_logger().warning(msg)
else:
raise RuntimeError(msg)
mle = CmdStanMLE(runset)
return mle
return CmdStanMLE.from_files(
csv_file=runset.csv_files[0],
config_file=runset.config_files[0],
stdout_file=runset.stdout_files[0],
converged=converged,
)

# pylint: disable=too-many-arguments
def sample(
Expand Down Expand Up @@ -945,7 +949,16 @@ def sample(
)
get_logger().warning(msg)

mcmc = CmdStanMCMC(runset)
mcmc = CmdStanMCMC.from_files(
csv_files=runset.csv_files,
config_files=runset.config_files,
metric_files=runset.metric_files or None,
stdout_files=runset.stdout_files,
diagnostic_files=runset.diagnostic_files or None,
profile_files=runset.profile_files or None,
chain_ids=runset.chain_ids,
sig_figs=runset._args.sig_figs,
)
return mcmc

def generate_quantities(
Expand Down Expand Up @@ -1040,7 +1053,13 @@ def generate_quantities(
),
):
fit_object = previous_fit
fit_csv_files = previous_fit.runset.csv_files
if isinstance(
previous_fit,
(CmdStanPathfinder, CmdStanLaplace, CmdStanMLE, CmdStanVB),
):
fit_csv_files = [previous_fit.csv_file]
else:
fit_csv_files = previous_fit.csv_files
elif isinstance(previous_fit, list):
if len(previous_fit) < 1:
raise ValueError(
Expand Down Expand Up @@ -1072,7 +1091,7 @@ def generate_quantities(
elif isinstance(fit_object, CmdStanMLE):
chains = 1
chain_ids = [1]
if fit_object._save_iterations:
if fit_object.config.method_config.save_iterations:
get_logger().warning(
'MLE contains saved iterations which will be used '
'to generate additional quantities of interest.'
Expand Down Expand Up @@ -1329,9 +1348,11 @@ def variational(
runset.get_err_msgs()
)
raise RuntimeError(msg)
# pylint: disable=invalid-name
vb = CmdStanVB(runset)
return vb
return CmdStanVB.from_files(
csv_file=runset.csv_files[0],
config_file=runset.config_files[0],
stdout_file=runset.stdout_files[0],
)

def pathfinder(
self,
Expand Down Expand Up @@ -1553,7 +1574,11 @@ def pathfinder(
' '.join(runset.cmd(0)), runset.get_err_msgs()
)
raise RuntimeError(msg)
return CmdStanPathfinder(runset)
return CmdStanPathfinder.from_files(
csv_file=runset.csv_files[0],
config_file=runset.config_files[0],
stdout_file=runset.stdout_files[0],
)

def log_prob(
self,
Expand Down Expand Up @@ -1739,24 +1764,20 @@ def laplace_sample(
else:
cmdstan_mode = mode

if cmdstan_mode.runset.method != Method.OPTIMIZE:
if not isinstance(cmdstan_mode, CmdStanMLE):
raise ValueError(
"Mode must be a CmdStanMLE or a path to an optimize CSV"
)

mode_jacobian = (
cmdstan_mode.runset._args.method_args.jacobian # type: ignore
)
mode_jacobian = cmdstan_mode.config.method_config.jacobian
if mode_jacobian != jacobian:
raise ValueError(
"Jacobian argument to optimize and laplace must match!\n"
f"Laplace was run with jacobian={jacobian},\n"
f"but optimize was run with jacobian={mode_jacobian}"
)

laplace_args = LaplaceArgs(
cmdstan_mode.runset.csv_files[0], draws, jacobian
)
laplace_args = LaplaceArgs(cmdstan_mode.csv_file, draws, jacobian)

with temp_single_json(data) as _data:
args = CmdStanArgs(
Expand All @@ -1780,7 +1801,12 @@ def laplace_sample(
timeout=timeout,
)
runset.raise_for_timeouts()
return CmdStanLaplace(runset, cmdstan_mode)
return CmdStanLaplace.from_files(
csv_file=runset.csv_files[0],
config_file=runset.config_files[0],
stdout_file=runset.stdout_files[0],
mode=cmdstan_mode,
)

def _run_cmdstan(
self,
Expand Down
Loading