From 87f66ac8229be768f1272a3999080dfba0012a01 Mon Sep 17 00:00:00 2001 From: Stefan Hackmann Date: Tue, 22 Sep 2026 13:48:46 +0000 Subject: [PATCH 1/8] docs: Document the plausibility chain and make the jitter relative --- README.md | 2 +- docs/api-reference.md | 14 +- notebooks/forecast_condition.ipynb | 323 ++++++++++++++-- src/jointfm_client/__init__.py | 4 + src/jointfm_client/adapters.py | 40 +- src/jointfm_client/client.py | 55 +++ src/jointfm_client/contract.py | 322 +++++++++++++++- .../fixtures/condition_log_prob_request.json | 51 +++ .../fixtures/condition_log_prob_response.json | 36 ++ tests/fixtures/forecast_log_prob_request.json | 41 ++ .../fixtures/forecast_log_prob_response.json | 33 ++ tests/test_fixture_compatibility.py | 13 + tests/test_log_prob_mode.py | 363 ++++++++++++++++++ 13 files changed, 1255 insertions(+), 42 deletions(-) create mode 100644 tests/fixtures/condition_log_prob_request.json create mode 100644 tests/fixtures/condition_log_prob_response.json create mode 100644 tests/fixtures/forecast_log_prob_request.json create mode 100644 tests/fixtures/forecast_log_prob_response.json create mode 100644 tests/test_log_prob_mode.py diff --git a/README.md b/README.md index 305ef18..50eab26 100644 --- a/README.md +++ b/README.md @@ -163,7 +163,7 @@ bootstrap_notebook(add_src_root=True) Run `task setup` first so VS Code can select the registered `Python (joint-client-python)` notebook kernel backed by this repository's `.venv`. -The bootstrap helper resolves the nearest src-layout Python project root, switches the working directory there, and prepends that project's local `src` tree during development. The examples cover hosted health checks, low-level JSON prediction, mean forecasts, sample forecasts, quantile forecasts, conditional forecasts, pandas/NumPy result conversion, and CSV forecast workflows. They use `.env.sample` placeholders and checked-in fixture payloads; no real tokens or deployment IDs are stored in notebooks. +The bootstrap helper resolves the nearest src-layout Python project root, switches the working directory there, and prepends that project's local `src` tree during development. The examples cover hosted health checks, low-level JSON prediction, mean forecasts, sample forecasts, quantile forecasts, conditional forecasts (one conditional read as a mean, as draws, as quantiles inside a bounded band, and as the log density of observed values, plus the ranking of candidate conditions), pandas/NumPy result conversion, and CSV forecast workflows. They use `.env.sample` placeholders and checked-in fixture payloads; no real tokens or deployment IDs are stored in notebooks. The current forecast request contract is: diff --git a/docs/api-reference.md b/docs/api-reference.md index 1f05497..e8de616 100644 --- a/docs/api-reference.md +++ b/docs/api-reference.md @@ -6,7 +6,7 @@ This reference covers the supported public Python surface exported by `jointfm_c | Name | Purpose | | --- | --- | -| `JointFMClient` | Synchronous client for hosted or local JointFM endpoints. Use `from_env()` for `.env` and `config.yaml` backed hosted settings, `health()` for consensus typed service metadata, `health_instances()` for per-deployment probe results and pooled sample topology, `predict(payload)` for low-level JSON prediction, `forecast(...)` for validated tabular forecasts, and the `forecast_mean(...)`, `forecast_samples(...)`, and `forecast_quantiles(...)` convenience methods for typed forecast results. Each forecast method accepts `condition=ConditionBlock(...)`, which switches the request to the `condition` query mode and, before anything is sent, checks that the deployment's health metadata advertises the mode and every condition kind the block uses. `health()` probes `GET /healthz` for local deployments and POSTs `{"request_type": "health"}` to `predict_url` for hosted DataRobot deployments because the DataRobot deployment gateway only proxies the unstructured prediction route. `feature_importance(...)` runs permutation feature importance: one baseline `forecast_samples` call plus one per shuffled feature column, returning a list of `{"feature", "mean", "distance"}` dicts, each holding that feature's absolute forecast-mean shift and centered squared 2-Wasserstein distance indexed by target and horizon. | +| `JointFMClient` | Synchronous client for hosted or local JointFM endpoints. Use `from_env()` for `.env` and `config.yaml` backed hosted settings, `health()` for consensus typed service metadata, `health_instances()` for per-deployment probe results and pooled sample topology, `predict(payload)` for low-level JSON prediction, `forecast(...)` for validated tabular forecasts, and the `forecast_mean(...)`, `forecast_samples(...)`, `forecast_quantiles(...)`, and `forecast_log_prob(...)` convenience methods for typed forecast results. `forecast_log_prob(...)` is the one that asks about values the caller already holds: it takes `query_rows`, one observed row per entry of `query_times` carrying every declared column, and returns their log density under the model's joint. Each forecast method accepts `condition=ConditionBlock(...)`, which switches the request to the `condition` query mode and, before anything is sent, checks that the deployment's health metadata advertises the mode and every condition kind the block uses. `health()` probes `GET /healthz` for local deployments and POSTs `{"request_type": "health"}` to `predict_url` for hosted DataRobot deployments because the DataRobot deployment gateway only proxies the unstructured prediction route. `feature_importance(...)` runs permutation feature importance: one baseline `forecast_samples` call plus one per shuffled feature column, returning a list of `{"feature", "mean", "distance"}` dicts, each holding that feature's absolute forecast-mean shift and centered squared 2-Wasserstein distance indexed by target and horizon. | `JointFMClient.from_env()` loads `config.yaml`, optional `.env` values, and process environment variables. `JointFMClient.health(cache=True)` caches health metadata only when requested. `JointFMClient.health_instances()` returns the same probe as a `HealthInstances` object: one `InstanceHealth` per configured deployment (including failures), `max_sample_count` as the sum of reachable caps (overall parallel capacity), and `topology` / `topology_label` grouping those caps (unavailable peers are listed but excluded from the sum and topology). Each endpoint's health payload describes only that endpoint; the client aggregates by calling each configured peer. `health()` still exposes the minimum reachable `max_sample_count`, which is the sample-batch cap used by forecast helpers. `JointFMClient.predict(payload)` requires `payload["model_version"]`; high-level forecast helpers resolve the configured model version when the caller does not pass one explicitly. When `forecast_samples(...)` requests an explicit `n_samples`, the client learns the deployment's `max_sample_count` from health metadata before the first prediction, splits oversized requests into capped prediction batches, and returns one merged `SampleForecastResult`. Clients configured without a reachable health route fall back to discovering the cap from the structured service error. @@ -17,7 +17,7 @@ This reference covers the supported public Python surface exported by `jointfm_c | `ColumnSpec` | Describes one modeled request column. Fields are `name`, `modality`, `role`, `nullable`, `vocabulary_size`, `level_count`, `mapping`, `lower_bound`, `upper_bound`, `time_value_kind`, `time_value_scale_seconds`, `time_value_use_local_normalized_time`, `time_value_calendar_id`, and `time_value_timezone`. | | `DataFrameSchema` | Describes tabular history layout. Fields are `columns`, `time_index_mode`, `time_column`, `time_scale_seconds`, `use_local_normalized_time`, `calendar_id`, and `timezone`. | | `ForecastRequestMetadata` | Holds `schema_version`, `model_version`, `query_mode` (`forecast` or `condition`), and `return_mode` for one forecast request. | -| `ForecastRequest` | Validated request object that combines metadata, schema, history rows, query times, requested columns, sample or quantile controls, `seed`, and an optional `condition` block, then emits a JSON-compatible payload with `to_payload()`. `query_mode="condition"` and a `condition` block must appear together, and the block is validated against the request's own schema and `query_times` before any round trip. | +| `ForecastRequest` | Validated request object that combines metadata, schema, history rows, query times, requested columns, sample or quantile controls, `seed`, an optional `condition` block, and the optional `query_rows` a scored request supplies, then emits a JSON-compatible payload with `to_payload()`. `query_mode="condition"` and a `condition` block must appear together, as must `return_mode="log_prob"` and `query_rows`, and both are validated against the request's own schema and `query_times` before any round trip. | | `EqualityCondition` | One column of the conditioned position pinned to a finite `value`. The pinned column leaves the read-out set, so it must not appear in `requested_columns`. | | `IntervalCondition` | One column of the conditioned position confined to `[lower, upper]`; `None` leaves a side open, at least one side must be bounded, and `lower < upper`. The column stays readable and the response describes its distribution inside the range. | | `ConditionBlock` | Every condition of one request: `query_time_index` into the request's `query_times` and a sequence of `EqualityCondition` or `IntervalCondition` values, at most one per column. `kinds`, `pinned_columns`, and `conditioned_columns` expose what the capability gate and the request validation need. | @@ -31,11 +31,13 @@ This reference covers the supported public Python surface exported by `jointfm_c | `StructuredError` | One structured JointFM service error with `code`, `message`, and optional `field`. | | `ForecastDiagnostics` | Response diagnostics containing `history_rows`, `horizon_count`, optional `seed`, and, on condition responses, `condition_draws` and an optional `interval_estimator`. | | `QuantileForecast` | One quantile surface with `quantile` and `values`; `to_numpy()` returns axis order `(horizon, column)`. | -| `ForecastOutputs` | Legacy nested view of parsed output arrays, including `query_times`, `requested_columns`, and exactly one of `mean`, `samples`, or `quantiles`. | +| `ForecastOutputs` | Legacy nested view of parsed output arrays, including `query_times`, `requested_columns`, and exactly one of `mean`, `samples`, `quantiles`, or `log_prob`. A response whose `return_mode` no parser handles is refused rather than read as another mode's shape. | +| `LogProbScores` | One `outputs.log_prob` block: per-horizon `values` and `nll_values`, plus the `total`, `mean`, `nll_total`, and `nll_mean` summaries. The summaries are redundant, so parsing recomputes them from `values` and rejects a payload that disagrees with itself. A log density is joint over the scored columns, so the block has no column axis. | | `ForecastResponse` | Shared base for parsed forecast results. It preserves schema, image, model, checkpoint, head, mode, query-time, requested-column, diagnostic, and error metadata, plus `plausibility` (`None` on forecast responses). | | `MeanForecastResult` | Parsed mean forecast. `to_numpy()` returns `(horizon, column)` and pandas helpers return tidy or wide frames. | | `SampleForecastResult` | Parsed sample forecasts. `to_numpy()` returns `(sample, horizon, column)` and pandas helpers return tidy or wide frames. | | `QuantileForecastResult` | Parsed quantile forecasts. `to_numpy()` returns `(quantile, horizon, column)`, `quantile_levels` exposes the ordered levels, and pandas helpers return tidy or wide frames. | +| `LogProbResult` | Parsed log densities of the values the request supplied; it carries no forecast values. `to_numpy()` returns `(horizon,)`, `to_pandas_tidy()` one row per scored position, and `to_pandas_wide()` the one-row request summary, since a log density has no column axis to widen. | ## Configuration Classes @@ -158,13 +160,14 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | `schema_version` | Yes | Must be `"v3"`. | | `model_version` | Yes | Exact deployed model version expected by the caller. | | `query_mode` | Yes | `"forecast"` for the unconditional forecast or `"condition"` for a conditional query at one future position. | -| `return_mode` | Yes | One of `"mean"`, `"samples"`, `"quantiles"`, or `"log_prob"`. The high-level `forecast_mean`, `forecast_samples`, and `forecast_quantiles` helpers cover the first three; `"log_prob"` is reachable through the low-level `predict(payload)` path. | +| `return_mode` | Yes | One of `"mean"`, `"samples"`, `"quantiles"`, or `"log_prob"`, covered by the `forecast_mean`, `forecast_samples`, `forecast_quantiles`, and `forecast_log_prob` helpers. | | `time_index_mode` | Yes | One of `"ordinal"`, `"continuous_float"`, or `"absolute_datetime"`. | | `columns` | Yes | Non-empty array of column descriptors for modeled columns. | | `history_rows` | Yes | Non-empty array of history row objects in the declared schema. | | `query_times` | Yes | Non-empty future forecast horizon values. Absolute datetimes are encoded timezone-stably. | | `time_column` | For absolute datetime, optional otherwise | Name of the history time column. It must not duplicate a modeled column name. | -| `requested_columns` | Optional | Output column names or integer indices. Duplicates are rejected. Defaults to all modeled columns. | +| `requested_columns` | Optional | Output column names or integer indices. Duplicates are rejected. Defaults to all modeled columns, minus the columns an equality condition pinned. Under `return_mode="log_prob"` an explicit list must name every readable column in schema order, because the score is joint over them. | +| `query_rows` | `log_prob` mode | Observed rows to score, one per entry of `query_times`, each carrying a value for every declared column. Required by `return_mode="log_prob"` and rejected for every other mode. A row that contradicts a condition — a pinned column given another value, or a bounded column outside its range — is refused by the service instead of scored. | | `n_samples` | Samples and quantiles controls | Positive sample count when sampling controls are needed. Oversized sample forecasts are batched automatically against the cap advertised in health metadata. | | `quantiles` | Quantiles mode | Quantile levels in `(0, 1)`, required for `return_mode="quantiles"`. | | `seed` | Optional | Integer random seed for reproducible stochastic outputs. | @@ -212,6 +215,7 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | `diagnostics.history_rows` | Number of history rows processed. | | `diagnostics.horizon_count` | Number of forecast horizon steps returned. | | `diagnostics.seed` | Optional seed used by the service. | +| `outputs.log_prob` | `log_prob` responses only: per-horizon `values` and `nll_values` with their `total`, `mean`, `nll_total`, and `nll_mean` summaries. | | `diagnostics.condition_draws` | Condition responses only: number of draws behind sampled outputs. Batched sample requests report the merged count. | | `diagnostics.interval_estimator` | Condition responses only, and only when more than one column carries an interval: `points` and `effective_sample_size` of the numerical region-probability estimate. | | `plausibility` | `null` on forecast responses. On condition responses an object with `equality_log_density` (log density of the pinned values) and `region_log_probability` (log probability of the interval region), each `null` when the request carried no condition of that kind. | diff --git a/notebooks/forecast_condition.ipynb b/notebooks/forecast_condition.ipynb index 521c81e..d3fc9f9 100644 --- a/notebooks/forecast_condition.ipynb +++ b/notebooks/forecast_condition.ipynb @@ -54,8 +54,12 @@ "- An **equality condition** pins a column to a value (`EqualityCondition`). The pinned column leaves the read-out set, because reading it back would only repeat the request.\n", "- An **interval condition** confines a column to a range whose bounds may be open on either side (`IntervalCondition`). The column stays readable: what comes back is its distribution inside the range.\n", "\n", + "The cell below asks for the same forecast twice, once without the block and once with it, because a conditional mean is only readable next to the unconditional one: what the what-if is worth is the shift between them, not the level of either.\n", + "\n", "Every column without a condition is a read-out column, and the response describes the conditional distribution of those columns at the conditioned position only, so `outputs.query_times` has exactly one entry however many `query_times` the request carried.\n", "\n", + "Every read-out the service serves reads that same conditional, and the sections below take it in turn: as a mean beside the unconditional one, as coherent joint draws, as quantiles inside a bounded band, and as the log density of values you already hold. The last two sections use the plausibility numbers to rank candidate what-ifs against each other.\n", + "\n", "Whether a deployment can condition depends on the checkpoint's head, so `/healthz` advertises `condition` in `supported_query_modes` and the kinds it answers in `supported_condition_kinds`. The client checks that advertisement before sending, and this notebook reads it explicitly so the check is visible." ] }, @@ -83,12 +87,16 @@ "HISTORY_PATH = Path(\"notebooks/history.csv\")\n", "FEATURE_COLUMNS = [\"equity_index_level\", \"treasury_10y_yield\", \"eur_usd_rate\"]\n", "TARGET_COLUMNS = [\"portfolio_nav\", \"realized_volatility\"]\n", + "PINNED_COLUMN = \"equity_index_level\"\n", "INPUT_STEPS = 100\n", "OUTPUT_HORIZONS = 10\n", "CONDITIONED_STEP = 0\n", "EQUITY_RALLY = 1.02\n", "EXPECTED_COLUMNS = FEATURE_COLUMNS + TARGET_COLUMNS\n", "QUERY_TIMES = list(range(INPUT_STEPS, INPUT_STEPS + OUTPUT_HORIZONS))\n", + "# Every column a request may read back: an equality condition removes its own\n", + "# column from the read-out set, and nothing else does.\n", + "READABLE_COLUMNS = [name for name in EXPECTED_COLUMNS if name != PINNED_COLUMN]\n", "\n", "history = pd.read_csv(HISTORY_PATH, dtype=float)\n", "if list(history.columns) != EXPECTED_COLUMNS:\n", @@ -115,14 +123,18 @@ " query_times_length=len(QUERY_TIMES),\n", ")\n", "\n", - "last_equity_level = float(history[\"equity_index_level\"].iloc[-1])\n", + "last_equity_level = float(history[PINNED_COLUMN].iloc[-1])\n", + "rallied_equity_level = last_equity_level * EQUITY_RALLY\n", "rally = ConditionBlock(\n", " query_time_index=CONDITIONED_STEP,\n", - " conditions=[\n", - " EqualityCondition(\n", - " column=\"equity_index_level\", value=last_equity_level * EQUITY_RALLY\n", - " )\n", - " ],\n", + " conditions=[EqualityCondition(column=PINNED_COLUMN, value=rallied_equity_level)],\n", + ")\n", + "baseline = client.forecast_mean(\n", + " history,\n", + " query_times=QUERY_TIMES,\n", + " requested_columns=plan.requested_columns,\n", + " columns=plan.columns,\n", + " seed=7,\n", ")\n", "result = client.forecast_mean(\n", " history,\n", @@ -145,7 +157,18 @@ " raise ValueError(\n", " f\"Expected {expected_forecast_rows} forecast rows, got {len(forecast)}\"\n", " )\n", - "forecast" + "\n", + "# The conditional answers one position, so joining on it keeps exactly the rows\n", + "# the two requests have in common.\n", + "comparison = baseline.to_pandas_tidy().merge(\n", + " forecast,\n", + " on=[\"query_time\", \"requested_column\"],\n", + " suffixes=(\"_unconditional\", \"_given_the_rally\"),\n", + ")\n", + "comparison[\"shift\"] = (\n", + " comparison[\"value_given_the_rally\"] - comparison[\"value_unconditional\"]\n", + ")\n", + "comparison" ] }, { @@ -157,27 +180,76 @@ }, "source": [ "## Reading the plausibility\n", - "`result.plausibility.equality_log_density` is the log density the model assigns to the pinned value before conditioning. It separates *the model is confident about NAV given this rally* from *the model finds a rally of this size absurd and is extrapolating*. The service reports the number and never refuses on it; comparing it across candidate pins, or against the density of a pin at the model's own unconditional mean, is the caller's decision." + "`result.plausibility.equality_log_density` is the log density the model assigns to the pinned value before conditioning. It separates *the model is confident about NAV given this rally* from *the model finds a rally of this size absurd and is extrapolating*. The service reports the number and never refuses on it; comparing it across candidate pins, or against the density of a pin at the model's own unconditional mean, is the caller's decision.\n", + "\n", + "It is a *density*, not a probability, and `exp()` of it is not one either: an exact value of a continuous column has probability zero, and a density carries the reciprocal units of the column it scores, so quoting the equity index in thousands of points instead of points shifts this number by `log(1000)` and can push its exponential above one. Only differences are unit-free, which is why the sentence above says compare rather than threshold: the gap between two pins on the same column is a log likelihood ratio, and that is what ranks candidate scenarios. When a pin covers several columns at once the number stays a single joint density over all of them — never a per-column value to be multiplied back together, which would assume the columns move independently." ] }, { "cell_type": "markdown", "id": "5", + "metadata": { + "id": "condition-draws-description", + "language": "markdown" + }, + "source": [ + "## Scenarios, not summaries\n", + "A mean answers *where*, and a band answers *how wide*, but a portfolio question usually needs whole futures: draws. Every row below is one coherent joint scenario, because the service draws the read-out columns together from the same conditional mixture — the NAV and the volatility in one row belong to each other, so a function of several columns can be evaluated row by row. A quantile table cannot answer that: the 90th percentile of NAV and the 90th percentile of volatility need not describe any single future.\n", + "\n", + "`diagnostics.condition_draws` reports how many draws stand behind the answer. Oversized sample requests are split into batches against the deployment's advertised cap and merged locally, and the merged count is what this field reports." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "6", + "metadata": { + "id": "condition-draws", + "language": "python" + }, + "outputs": [], + "source": [ + "N_SAMPLES = 8\n", + "\n", + "scenarios = client.forecast_samples(\n", + " history,\n", + " query_times=QUERY_TIMES,\n", + " requested_columns=plan.requested_columns,\n", + " columns=plan.columns,\n", + " n_samples=N_SAMPLES,\n", + " seed=7,\n", + " condition=rally,\n", + ")\n", + "if scenarios.query_times != (QUERY_TIMES[CONDITIONED_STEP],):\n", + " raise ValueError(\n", + " f\"Expected the conditioned position alone, got {scenarios.query_times!r}\"\n", + " )\n", + "print(\"draws behind the answer:\", scenarios.diagnostics.condition_draws)\n", + "scenarios.to_pandas_wide()" + ] + }, + { + "cell_type": "markdown", + "id": "7", "metadata": { "id": "condition-interval-description", "language": "markdown" }, "source": [ "## Interval condition\n", - "Now confine the 10-year yield to a band around its last observed value instead of pinning it, and pin the equity index at the same time: a request may mix both kinds across the columns of one position. The response then carries `region_log_probability`, the log probability the model gives the yield band, and the yield column itself stays readable because its distribution inside the band is a genuine answer.\n", + "Now confine the 10-year yield to a band around its last observed value instead of pinning it, and pin the equity index at the same time: a request may mix both kinds across the columns of one position. The response then carries `region_log_probability`, and the yield column itself stays readable because its distribution inside the band is a genuine answer.\n", + "\n", + "That number is the log probability of the band **given the pinned equity index**, not the band's own probability, because the two condition kinds compose in a fixed order and the band is measured on the distribution the pin has already reduced. The two numbers therefore chain rather than describe separate things, and adding them gives the plausibility of the whole request — `log(density(pin) * P(band | pin))` — with no independence assumed anywhere. A request carrying only interval conditions has nothing to combine, and its `region_log_probability` alone is the joint probability of everything it asked about.\n", + "\n", + "With one interval column, as here, the region probability is exact — one difference of distribution functions per mixture component — and `diagnostics.interval_estimator` stays empty, because there is no estimate to characterize. Bounding a second column makes the region a box with no closed form, which the service estimates numerically and then reports the accounting for — the next section does exactly that.\n", "\n", - "With one interval column the region probability is exact. When several columns carry intervals the service estimates the probability of the box numerically and reports the accounting in `diagnostics.interval_estimator`, so the caller can judge the estimate." + "A band is also where the quantile read-out earns its place over the mean: what the banded column comes back with is a distribution inside its own range, so the cell asserts every quantile of it lands there." ] }, { "cell_type": "code", "execution_count": null, - "id": "6", + "id": "8", "metadata": { "id": "condition-interval-example", "language": "python" @@ -186,27 +258,31 @@ "source": [ "from jointfm_client import IntervalCondition\n", "\n", + "QUANTILES = [0.1, 0.5, 0.9]\n", "YIELD_BAND_HALF_WIDTH = 0.002\n", + "# A quantile of the truncated draws can sit on the band edge, and the round trip\n", + "# through the response's float encoding may move it by an ulp.\n", + "BAND_TOLERANCE = 1e-9\n", "\n", "last_yield = float(history[\"treasury_10y_yield\"].iloc[-1])\n", + "yield_lower = last_yield - YIELD_BAND_HALF_WIDTH\n", + "yield_upper = last_yield + YIELD_BAND_HALF_WIDTH\n", + "yield_band = IntervalCondition(\n", + " column=\"treasury_10y_yield\", lower=yield_lower, upper=yield_upper\n", + ")\n", "rally_with_yield_band = ConditionBlock(\n", " query_time_index=CONDITIONED_STEP,\n", " conditions=[\n", - " EqualityCondition(\n", - " column=\"equity_index_level\", value=last_equity_level * EQUITY_RALLY\n", - " ),\n", - " IntervalCondition(\n", - " column=\"treasury_10y_yield\",\n", - " lower=last_yield - YIELD_BAND_HALF_WIDTH,\n", - " upper=last_yield + YIELD_BAND_HALF_WIDTH,\n", - " ),\n", + " EqualityCondition(column=PINNED_COLUMN, value=rallied_equity_level),\n", + " yield_band,\n", " ],\n", ")\n", - "banded = client.forecast_mean(\n", + "banded = client.forecast_quantiles(\n", " history,\n", " query_times=QUERY_TIMES,\n", " requested_columns=[\"treasury_10y_yield\", *plan.requested_columns],\n", " columns=plan.columns,\n", + " quantiles=QUANTILES,\n", " seed=7,\n", " condition=rally_with_yield_band,\n", ")\n", @@ -214,16 +290,215 @@ " raise ValueError(\"A condition response must carry its plausibility block\")\n", "print(\"log density of the pinned value:\", banded.plausibility.equality_log_density)\n", "print(\"log probability of the yield band:\", banded.plausibility.region_log_probability)\n", - "print(\"interval estimator accounting:\", banded.diagnostics.interval_estimator)\n", - "banded.to_pandas_tidy()" + "if banded.diagnostics.interval_estimator is not None:\n", + " raise ValueError(\"A single interval column is exact and reports no estimator\")\n", + "band = banded.to_pandas_wide()\n", + "outside_band = band[\n", + " (band[\"treasury_10y_yield\"] < yield_lower - BAND_TOLERANCE)\n", + " | (band[\"treasury_10y_yield\"] > yield_upper + BAND_TOLERANCE)\n", + "]\n", + "if not outside_band.empty:\n", + " raise ValueError(f\"A banded column left its own interval:\\n{outside_band}\")\n", + "band" + ] + }, + { + "cell_type": "markdown", + "id": "9", + "metadata": { + "id": "condition-box-description", + "language": "markdown" + }, + "source": [ + "## When the region has to be estimated\n", + "Bounding the FX rate as well turns the region into a box. That has no closed form, so the service estimates its probability with a quasi-random walk and only then reports the accounting in `diagnostics.interval_estimator`: `points` is how many quasi-random points went into the estimate and `effective_sample_size` is Kish's effective sample size of the weights behind them. A small effective sample size against a large point count means a deep-tail box where few points carry the answer. The cell asserts both sides of that rule — absent for the single band above, present here.\n", + "\n", + "Drawing from a boxed conditional also shows what the box does to the draws. Each row pairs one truncated draw of the bounded columns with a read-out drawn from *that* draw's conditional, so the rows are exact draws from the joint conditional rather than separately summarized margins, and every one of them lands inside both bands." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "10", + "metadata": { + "id": "condition-box", + "language": "python" + }, + "outputs": [], + "source": [ + "FX_BAND_HALF_WIDTH = 0.004\n", + "\n", + "last_fx_rate = float(history[\"eur_usd_rate\"].iloc[-1])\n", + "fx_lower = last_fx_rate - FX_BAND_HALF_WIDTH\n", + "fx_upper = last_fx_rate + FX_BAND_HALF_WIDTH\n", + "fx_band = IntervalCondition(column=\"eur_usd_rate\", lower=fx_lower, upper=fx_upper)\n", + "rally_with_two_bands = ConditionBlock(\n", + " query_time_index=CONDITIONED_STEP,\n", + " conditions=[\n", + " EqualityCondition(column=PINNED_COLUMN, value=rallied_equity_level),\n", + " yield_band,\n", + " fx_band,\n", + " ],\n", + ")\n", + "boxed = client.forecast_samples(\n", + " history,\n", + " query_times=QUERY_TIMES,\n", + " requested_columns=[\"treasury_10y_yield\", \"eur_usd_rate\", *plan.requested_columns],\n", + " columns=plan.columns,\n", + " n_samples=N_SAMPLES,\n", + " seed=7,\n", + " condition=rally_with_two_bands,\n", + ")\n", + "if boxed.plausibility is None:\n", + " raise ValueError(\"A condition response must carry its plausibility block\")\n", + "estimator = boxed.diagnostics.interval_estimator\n", + "if estimator is None:\n", + " raise ValueError(\n", + " \"A multi-column region is estimated and must report its accounting\"\n", + " )\n", + "print(\n", + " \"log probability of the box given the pin:\",\n", + " boxed.plausibility.region_log_probability,\n", + ")\n", + "print(\"estimator points:\", estimator.points)\n", + "print(\"effective sample size:\", estimator.effective_sample_size)\n", + "box_draws = boxed.to_pandas_wide()\n", + "outside_box = box_draws[\n", + " (box_draws[\"treasury_10y_yield\"] < yield_lower - BAND_TOLERANCE)\n", + " | (box_draws[\"treasury_10y_yield\"] > yield_upper + BAND_TOLERANCE)\n", + " | (box_draws[\"eur_usd_rate\"] < fx_lower - BAND_TOLERANCE)\n", + " | (box_draws[\"eur_usd_rate\"] > fx_upper + BAND_TOLERANCE)\n", + "]\n", + "if not outside_box.empty:\n", + " raise ValueError(f\"A draw left the box it was conditioned into:\\n{outside_box}\")\n", + "box_draws" + ] + }, + { + "cell_type": "markdown", + "id": "11", + "metadata": { + "id": "condition-ranking-description", + "language": "markdown" + }, + "source": [ + "## Ranking candidate scenarios\n", + "The plausibility section above ended on the rule that only differences of `equality_log_density` are unit-free. This is what that buys: the gap between two pins on the same column is a log likelihood ratio, so candidate what-ifs can be ordered even though no single one of them can be thresholded.\n", + "\n", + "The table asks the model for the conditional forecast under several rallies and reports each pin's plausibility relative to the most plausible one, so a scenario the model finds far-fetched is visible next to the answer it produced. The mean is the right read-out here precisely because the question is not about one scenario's shape: it needs one comparable number per candidate." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "12", + "metadata": { + "id": "condition-ranking", + "language": "python" + }, + "outputs": [], + "source": [ + "CANDIDATE_RALLIES = [0.94, 0.98, 1.00, 1.02, 1.06]\n", + "\n", + "rankings = []\n", + "for candidate in CANDIDATE_RALLIES:\n", + " candidate_level = last_equity_level * candidate\n", + " candidate_result = client.forecast_mean(\n", + " history,\n", + " query_times=QUERY_TIMES,\n", + " requested_columns=plan.requested_columns,\n", + " columns=plan.columns,\n", + " seed=7,\n", + " condition=ConditionBlock(\n", + " query_time_index=CONDITIONED_STEP,\n", + " conditions=[EqualityCondition(column=PINNED_COLUMN, value=candidate_level)],\n", + " ),\n", + " )\n", + " if candidate_result.plausibility is None:\n", + " raise ValueError(\"A condition response must carry its plausibility block\")\n", + " conditional_mean = candidate_result.to_pandas_wide().iloc[0]\n", + " rankings.append(\n", + " {\n", + " \"rally\": candidate,\n", + " PINNED_COLUMN: candidate_level,\n", + " \"equality_log_density\": candidate_result.plausibility.equality_log_density,\n", + " **{column: conditional_mean[column] for column in plan.requested_columns},\n", + " }\n", + " )\n", + "\n", + "ranking = pd.DataFrame.from_records(rankings)\n", + "ranking[\"log_ratio_vs_best\"] = (\n", + " ranking[\"equality_log_density\"] - ranking[\"equality_log_density\"].max()\n", + ")\n", + "ranking.sort_values(\"log_ratio_vs_best\", ascending=False, ignore_index=True)" + ] + }, + { + "cell_type": "markdown", + "id": "13", + "metadata": { + "id": "condition-score-description", + "language": "markdown" + }, + "source": [ + "## Scoring an outcome you already have\n", + "Every read-out so far answers *what does the model expect*. `log_prob` asks the opposite — how plausible are values I already hold — and is the only return mode where the caller supplies the future instead of receiving it. Under a condition it scores those values against the conditional, which is how two candidate outcomes of the same what-if are compared.\n", + "\n", + "Three rules follow from scoring a joint rather than a projection of one, and the client enforces all three before anything is sent. `query_rows` carries one observed row per entry of `query_times`, each with a value for every declared column. `requested_columns` must name every *readable* column in schema order — a narrower projection would score a different distribution than the caller means to ask about — and the pinned column is the one exception, since conditioning fixed its value. The deployment refuses a row that contradicts the condition, a pinned column given another value or a banded column outside its range, instead of scoring it.\n", + "\n", + "The scores below are log densities of whole rows, so the unit caveat from the plausibility section applies unchanged: read their difference, not their level." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "14", + "metadata": { + "id": "condition-score", + "language": "python" + }, + "outputs": [], + "source": [ + "SCORED_QUERY_TIMES = QUERY_TIMES[: CONDITIONED_STEP + 1]\n", + "STRESS_NAV_DROP = 0.9\n", + "\n", + "last_row = history.iloc[-1]\n", + "continuation = {\n", + " PINNED_COLUMN: rallied_equity_level,\n", + " \"treasury_10y_yield\": last_yield,\n", + " \"eur_usd_rate\": last_fx_rate,\n", + " \"portfolio_nav\": float(last_row[\"portfolio_nav\"]),\n", + " \"realized_volatility\": float(last_row[\"realized_volatility\"]),\n", + "}\n", + "stressed = {\n", + " **continuation,\n", + " \"portfolio_nav\": continuation[\"portfolio_nav\"] * STRESS_NAV_DROP,\n", + "}\n", + "\n", + "scores = {}\n", + "for label, outcome in {\"continuation\": continuation, \"stressed\": stressed}.items():\n", + " scored = client.forecast_log_prob(\n", + " history,\n", + " query_times=SCORED_QUERY_TIMES,\n", + " query_rows=pd.DataFrame([outcome], columns=EXPECTED_COLUMNS),\n", + " requested_columns=READABLE_COLUMNS,\n", + " columns=plan.columns,\n", + " seed=7,\n", + " condition=rally,\n", + " )\n", + " scores[label] = scored.log_prob.total\n", + "\n", + "print(\"log density of the continuation:\", scores[\"continuation\"])\n", + "print(\"log density of the stressed outcome:\", scores[\"stressed\"])\n", + "print(\"log likelihood ratio:\", scores[\"continuation\"] - scores[\"stressed\"])" ] } ], "metadata": { "kernelspec": { - "display_name": "joint-client-python (3.13.3)", + "display_name": "Python (joint-client-python)", "language": "python", - "name": "python3" + "name": "joint-client-python" }, "language_info": { "codemirror_mode": { diff --git a/src/jointfm_client/__init__.py b/src/jointfm_client/__init__.py index 5972c76..ffe9a94 100644 --- a/src/jointfm_client/__init__.py +++ b/src/jointfm_client/__init__.py @@ -64,6 +64,8 @@ IMPORT_NAMESPACE, LOCAL_HEALTH_ROUTE, LOCAL_PREDICT_ROUTE, + LogProbResult, + LogProbScores, PACKAGE_VERSION, PREDICT_REQUEST_TYPE, SCHEMA_VERSION, @@ -214,6 +216,8 @@ "JSONTransport", "LOCAL_HEALTH_ROUTE", "LOCAL_PREDICT_ROUTE", + "LogProbResult", + "LogProbScores", "QuantileForecast", "QuantileForecastResult", "PathConfig", diff --git a/src/jointfm_client/adapters.py b/src/jointfm_client/adapters.py index 988cb22..d1aa57d 100644 --- a/src/jointfm_client/adapters.py +++ b/src/jointfm_client/adapters.py @@ -171,8 +171,15 @@ def infer_column_specs_from_dataframe( def dataframe_to_history_rows( frame: Any, schema: DataFrameSchema, + *, + field: str = "history_rows", ) -> list[dict[str, Any]]: - """Convert a pandas ``DataFrame`` into ordered JointFM ``history_rows``.""" + """Convert a pandas ``DataFrame`` into ordered JointFM row payloads. + + ``field`` names the payload array being built, so a bad value in the + observed rows of a scored request reports ``query_rows[...]`` rather than + pointing the caller at their history. + """ pandas_module = _require_pandas() numpy_module = _require_numpy() if not isinstance(frame, pandas_module.DataFrame): @@ -198,7 +205,7 @@ def dataframe_to_history_rows( row_payload[column_name] = _time_index_value_to_json( value, time_index_mode=schema.time_index_mode, - field=f"history_rows[{row_index}].{column_name}", + field=f"{field}[{row_index}].{column_name}", pandas_module=pandas_module, numpy_module=numpy_module, ) @@ -207,7 +214,7 @@ def dataframe_to_history_rows( row_payload[column_name] = _column_value_to_json( value, column_spec=column_spec, - field=f"history_rows[{row_index}].{column_name}", + field=f"{field}[{row_index}].{column_name}", pandas_module=pandas_module, numpy_module=numpy_module, ) @@ -298,12 +305,18 @@ def build_forecast_payload_from_dataframe( nullable_columns: Sequence[str] | None = None, bounds: ColumnBounds | None = None, condition: ConditionBlock | None = None, + query_rows: Any | None = None, ) -> dict[str, Any]: """Build a validated forecast payload from a pandas ``DataFrame``. Passing ``condition`` makes this a conditioning request: the payload then carries ``query_mode='condition'`` and the block, and the deployment answers the conditional at the one future position the block names. + + ``query_rows`` carries the *observed* values at ``query_times`` that + ``return_mode='log_prob'`` scores — a ``DataFrame`` shaped like ``frame``, + or a sequence of row mappings. The two go together: neither is accepted + without the other. """ column_specs = ( infer_column_specs_from_dataframe( @@ -337,6 +350,9 @@ def build_forecast_payload_from_dataframe( timezone=timezone, ) history_rows = dataframe_to_history_rows(frame, schema) + query_row_payloads = ( + None if query_rows is None else _query_rows_payload(query_rows, schema) + ) normalized_query_times = validate_forecast_horizon( _history_times_from_dataframe(frame, schema), query_times, @@ -355,9 +371,27 @@ def build_forecast_payload_from_dataframe( schema_version=schema_version, query_mode="forecast" if condition is None else "condition", condition=condition, + query_rows=query_row_payloads, ) +def _query_rows_payload( + query_rows: Any, + schema: DataFrameSchema, +) -> list[dict[str, Any]]: + """Normalize observed scoring rows into row payloads. + + A ``DataFrame`` goes through the same column encoding as the history, so a + categorical label or a timestamp is spelled identically in both arrays; + anything else is already a sequence of row mappings and is passed through + for ``ForecastRequest`` to validate. + """ + pandas_module = _require_pandas() + if isinstance(query_rows, pandas_module.DataFrame): + return dataframe_to_history_rows(query_rows, schema, field="query_rows") + return [dict(row) for row in query_rows] + + def build_forecast_payload_from_arrays( values: Any, *, diff --git a/src/jointfm_client/client.py b/src/jointfm_client/client.py index 2b969a8..c15b1f9 100644 --- a/src/jointfm_client/client.py +++ b/src/jointfm_client/client.py @@ -64,6 +64,7 @@ from jointfm_client.contract import ( ForecastDiagnostics, ForecastResponse, + LogProbResult, MeanForecastResult, QuantileForecastResult, SampleForecastResult, @@ -313,6 +314,7 @@ def forecast( bounds: Mapping[str, tuple[float | int | None, float | int | None]] | None = None, condition: ConditionBlock | None = None, + query_rows: Any | None = None, ) -> ForecastResponse: """Build and submit a forecast request from tabular history inputs. @@ -320,6 +322,9 @@ def forecast( future position the block names, instead of the unconditional forecast. The deployment's advertised capability is checked first, so a deployment that cannot condition is refused here rather than after a round trip. + + ``query_rows`` carries the observed values at ``query_times`` that + ``return_mode='log_prob'`` scores, in the same shape as ``history``. """ self._require_predict_url("forecast") if condition is not None: @@ -348,6 +353,7 @@ def forecast( calendar_id=calendar_id, timezone=timezone, condition=condition, + query_rows=query_rows, ) else: payload = build_forecast_payload_from_dataframe( @@ -382,6 +388,7 @@ def forecast( nullable_columns=nullable_columns, bounds=bounds, condition=condition, + query_rows=query_rows, ) sample_cap = self._resolve_sample_batch_cap(payload) if sample_cap is not None: @@ -512,6 +519,52 @@ def forecast_quantiles( ), ) + def forecast_log_prob( + self, + history: Any, + *, + query_times: Sequence[Any], + query_rows: Any, + schema: DataFrameSchema | None = None, + time_index_mode: TimeIndexMode = "ordinal", + columns: Sequence[ColumnSpec] | None = None, + time_column: str | None = None, + requested_columns: Sequence[str | int] | None = None, + model_version: str | None = None, + seed: int | None = None, + condition: ConditionBlock | None = None, + ) -> LogProbResult: + """Score observed future values through the shared forecast validation path. + + This is the one return mode that answers a question about values the + caller already has: ``query_rows`` holds the observed row at each entry + of ``query_times``, and the result is the model's log density of those + values under its joint at that position. It therefore scores the whole + joint, and ``requested_columns`` — when given at all — must name every + readable column in schema order. + + With ``condition`` the score is taken under the conditional at the one + future position the block names, and the service refuses rows that + contradict the condition instead of scoring them; see :meth:`forecast`. + """ + return cast( + LogProbResult, + self.forecast( + history, + query_times=query_times, + schema=schema, + time_index_mode=time_index_mode, + columns=columns, + time_column=time_column, + requested_columns=requested_columns, + return_mode="log_prob", + model_version=model_version, + seed=seed, + condition=condition, + query_rows=query_rows, + ), + ) + def feature_importance( self, history: Any, @@ -668,6 +721,7 @@ def _forecast_payload_from_rows( calendar_id: str, timezone: str | None, condition: ConditionBlock | None = None, + query_rows: Any | None = None, ) -> dict[str, Any]: if schema is None: if columns is None: @@ -696,6 +750,7 @@ def _forecast_payload_from_rows( schema_version=schema_version, query_mode="forecast" if condition is None else "condition", condition=condition, + query_rows=query_rows, ) def _resolve_sample_batch_cap(self, payload: Mapping[str, Any]) -> int | None: diff --git a/src/jointfm_client/contract.py b/src/jointfm_client/contract.py index 1675796..208a1d9 100644 --- a/src/jointfm_client/contract.py +++ b/src/jointfm_client/contract.py @@ -38,6 +38,9 @@ PACKAGE_VERSION: Final = importlib.metadata.version(DISTRIBUTION_NAME) SCHEMA_VERSION: Final = "v3" +# The service sums and averages its log densities in the model's own precision, +# so its summaries differ from a recomputation in the last few digits. +DERIVED_SCORE_TOLERANCE: Final = 1e-6 DATAROBOT_UNSTRUCTURED_PREDICTION_ROUTE_TEMPLATE: Final = ( "deployments/{deployment_id}/predictionsUnstructured" ) @@ -503,6 +506,7 @@ class ForecastRequest: seed: int | None = None query_row_ids: Sequence[int] | None = None condition: ConditionBlock | None = None + query_rows: Sequence[Mapping[str, Any]] | None = None def __post_init__(self) -> None: """Validate payload controls and JSON-facing request arrays.""" @@ -539,6 +543,14 @@ def __post_init__(self) -> None: "quantiles may be provided only when return_mode='quantiles'" ) + if (self.metadata.return_mode == "log_prob") != (self.query_rows is not None): + raise ValueError( + "return_mode='log_prob' and query_rows go together: scoring needs " + "the observed values, and no other return mode reads them" + ) + if self.query_rows is not None: + self._validate_query_rows() + def _validate_condition(self) -> None: """Check the condition block against this request's schema and horizons. @@ -583,11 +595,73 @@ def _validate_condition(self) -> None: "the request supplied" ) + def _validate_query_rows(self) -> None: + """Check the observed rows a log-density request scores. + + The service scores the whole joint at each future position, so a row + must carry every declared column and ``requested_columns`` must name + every readable one in schema order: a narrower projection would score a + different distribution than the caller believes it asked about. An + equality condition is the one exception, since it fixes its column's + value and the conditioning contract forbids reading it back. + + Only two per-value failures are decidable without the deployment's own + encoding and are checked here — a missing value and a non-finite number. + A categorical label passes, because the service maps it through the + column's declared mapping before scoring it. + """ + rows = _require_sequence(self.query_rows, field="query_rows") + query_times = _require_sequence(self.query_times, field="query_times") + if len(rows) != len(query_times): + raise ValueError( + f"query_rows must carry one row per query time: got {len(rows)} " + f"rows for {len(query_times)} query_times" + ) + declared = [column.name for column in self.schema.columns] + required = ( + declared + if self.schema.time_column is None + else [ + *declared, + self.schema.time_column, + ] + ) + for index, row in enumerate(rows): + row_mapping = _require_mapping(row, field=f"query_rows[{index}]") + missing = [name for name in required if name not in row_mapping] + if missing: + raise ValueError( + f"query_rows[{index}] is missing declared columns: {missing}" + ) + for name in declared: + _require_scored_value( + row_mapping[name], field=f"query_rows[{index}].{name}" + ) + + requested = _resolve_requested_columns( + self.schema.columns, self.requested_columns + ) + if requested is None: + return + pinned = set(() if self.condition is None else self.condition.pinned_columns) + readable = [name for name in declared if name not in pinned] + if requested != readable: + excused = ( + f", the pinned columns {sorted(pinned)} excepted" if pinned else "" + ) + raise ValueError( + "return_mode='log_prob' scores the whole joint, so " + "requested_columns must list every declared column in schema " + f"order{excused}: expected {readable}, got {requested}" + ) + def to_payload(self) -> dict[str, Any]: """Return a JSON-compatible forecast request without mutating inputs.""" payload = self.metadata.to_payload() payload.update(self.schema.to_payload()) - payload["history_rows"] = _serialize_history_rows(self.history_rows) + payload["history_rows"] = _serialize_rows( + self.history_rows, field="history_rows" + ) payload["query_times"] = _serialize_query_times( self.query_times, time_index_mode=self.schema.time_index_mode, @@ -606,6 +680,8 @@ def to_payload(self) -> dict[str, Any]: payload["seed"] = self.seed if self.condition is not None: payload["condition"] = self.condition.to_payload() + if self.query_rows is not None: + payload["query_rows"] = _serialize_rows(self.query_rows, field="query_rows") return payload @@ -822,6 +898,28 @@ class ConditionPlausibility: space, and each is ``None`` when the request carried no condition of that kind. The service reports them and never refuses on them, so acting on them is the caller's decision. + + Each number is joint over its whole condition block rather than one value + per column: ``equality_log_density`` covers every pinned value at once, and + ``region_log_probability`` covers the whole box at once. Do not rebuild + either by multiplying per-column numbers, which would assume the columns are + independent and so discard what a joint model is for. + + ``region_log_probability`` is measured *after* the pinned columns are + applied, so on a request carrying both kinds it is the log probability of + the box **given the pins**, not the box's own. The two therefore chain, and + their sum is the plausibility of the whole condition set:: + + equality_log_density + region_log_probability + == log( density(pins) * P(box | pins) ) + + That sum is a density times a probability. Its ``exp`` is not a + probability -- it carries the reciprocal units of the pinned columns and can + exceed one -- so compare it, never threshold it: the difference between two + scenarios pinning the same columns is a log likelihood ratio and is + unit-free. For a request carrying only interval conditions there is nothing + to combine and ``region_log_probability`` alone is the joint probability of + the conditioned event. """ equality_log_density: float | None = None @@ -903,6 +1001,86 @@ def to_numpy(self) -> Any: return numpy_module.asarray(self.values, dtype=float) +@dataclass(frozen=True, slots=True) +class LogProbScores: + """Per-horizon log densities of one scored request, with its summaries. + + The service reports the same numbers several times over: ``values`` holds + the log density of each scored future position, ``nll_values`` their + negatives, and the four scalars the sum and the mean of each. Parsing + recomputes every redundant field from ``values`` instead of trusting it, so + a truncated or mismatched payload fails here rather than reading as a score. + + A log density is joint over the columns that were scored, so the block has + no column axis however wide the request was: one number per horizon. + """ + + values: tuple[float, ...] + nll_values: tuple[float, ...] + total: float + mean: float + nll_total: float + nll_mean: float + + @classmethod + def from_payload( + cls, + payload: Mapping[str, Any], + *, + expected_horizon_count: int | None = None, + ) -> Self: + """Parse one ``outputs.log_prob`` block and check its derived fields.""" + values = _require_float_vector( + payload.get("values"), + field="outputs.log_prob.values", + expected_length=expected_horizon_count, + ) + nll_values = _require_float_vector( + payload.get("nll_values"), + field="outputs.log_prob.nll_values", + expected_length=len(values), + ) + for index, (score, negated) in enumerate(zip(values, nll_values, strict=True)): + _require_derived_score( + negated, + expected=-score, + field=f"outputs.log_prob.nll_values[{index}]", + ) + summed = math.fsum(values) + total = _require_derived_score( + _require_output_float(payload.get("total"), field="outputs.log_prob.total"), + expected=summed, + field="outputs.log_prob.total", + ) + mean = _require_derived_score( + _require_output_float(payload.get("mean"), field="outputs.log_prob.mean"), + expected=summed / len(values), + field="outputs.log_prob.mean", + ) + nll_total = _require_derived_score( + _require_output_float( + payload.get("nll_total"), field="outputs.log_prob.nll_total" + ), + expected=-total, + field="outputs.log_prob.nll_total", + ) + nll_mean = _require_derived_score( + _require_output_float( + payload.get("nll_mean"), field="outputs.log_prob.nll_mean" + ), + expected=-mean, + field="outputs.log_prob.nll_mean", + ) + return cls( + values=values, + nll_values=nll_values, + total=total, + mean=mean, + nll_total=nll_total, + nll_mean=nll_mean, + ) + + @dataclass(frozen=True, slots=True) class ForecastOutputs: """Legacy-shaped forecast output arrays for one parsed forecast result.""" @@ -912,6 +1090,7 @@ class ForecastOutputs: mean: tuple[tuple[float, ...], ...] | None = None samples: tuple[tuple[tuple[float, ...], ...], ...] | None = None quantiles: tuple[QuantileForecast, ...] | None = None + log_prob: LogProbScores | None = None @classmethod def from_payload( @@ -948,6 +1127,7 @@ def from_payload( mean: tuple[tuple[float, ...], ...] | None = None samples: tuple[tuple[tuple[float, ...], ...], ...] | None = None quantiles: tuple[QuantileForecast, ...] | None = None + log_prob: LogProbScores | None = None if return_mode == "mean": mean = _require_float_matrix( @@ -958,6 +1138,7 @@ def from_payload( ) _require_none(payload.get("samples"), field="outputs.samples") _require_none(payload.get("quantiles"), field="outputs.quantiles") + _require_none(payload.get("log_prob"), field="outputs.log_prob") elif return_mode == "samples": samples = _require_float_tensor3( payload.get("samples"), @@ -968,7 +1149,8 @@ def from_payload( ) _require_none(payload.get("mean"), field="outputs.mean") _require_none(payload.get("quantiles"), field="outputs.quantiles") - else: + _require_none(payload.get("log_prob"), field="outputs.log_prob") + elif return_mode == "quantiles": quantiles = _require_quantile_forecasts( payload.get("quantiles"), field="outputs.quantiles", @@ -978,6 +1160,21 @@ def from_payload( ) _require_none(payload.get("mean"), field="outputs.mean") _require_none(payload.get("samples"), field="outputs.samples") + _require_none(payload.get("log_prob"), field="outputs.log_prob") + elif return_mode == "log_prob": + log_prob = LogProbScores.from_payload( + _require_mapping(payload.get("log_prob"), field="outputs.log_prob"), + expected_horizon_count=horizon_count, + ) + _require_none(payload.get("mean"), field="outputs.mean") + _require_none(payload.get("samples"), field="outputs.samples") + _require_none(payload.get("quantiles"), field="outputs.quantiles") + else: + # A mode nobody parses must fail here rather than fall through to + # another mode's branch and be read as that mode's shape. + raise ValueError( + f"outputs cannot be parsed for return_mode {return_mode!r}" + ) return cls( query_times=query_times, @@ -985,6 +1182,7 @@ def from_payload( mean=mean, samples=samples, quantiles=quantiles, + log_prob=log_prob, ) @@ -1168,8 +1366,11 @@ def from_payload( if return_mode == "samples": assert outputs.samples is not None return SampleForecastResult(samples=outputs.samples, **shared_fields) - assert outputs.quantiles is not None - return QuantileForecastResult(quantiles=outputs.quantiles, **shared_fields) + if return_mode == "quantiles": + assert outputs.quantiles is not None + return QuantileForecastResult(quantiles=outputs.quantiles, **shared_fields) + assert outputs.log_prob is not None + return LogProbResult(log_prob=outputs.log_prob, **shared_fields) @dataclass(frozen=True, slots=True) @@ -1370,6 +1571,66 @@ def to_pandas_wide(self) -> Any: return pandas_module.DataFrame.from_records(rows) +@dataclass(frozen=True, slots=True) +class LogProbResult(ForecastResponse): + """Parsed log densities of observed values, with shared metadata. + + This result carries no forecast values, because a scored request supplies + them itself as ``query_rows``: what comes back is how plausible the model + finds them. Each score is joint over the scored columns, so the conversions + have a horizon axis and no column axis. + """ + + log_prob: LogProbScores + + @property + def outputs(self) -> ForecastOutputs: + """Return the legacy nested outputs view for scored requests.""" + return ForecastOutputs( + query_times=self.query_times, + requested_columns=self.requested_columns, + log_prob=self.log_prob, + ) + + def to_numpy(self) -> Any: + """Return NumPy log densities with axis order ``(horizon,)``.""" + numpy_module = _require_numpy_module() + return numpy_module.asarray(self.log_prob.values, dtype=float) + + def to_pandas_tidy(self) -> Any: + """Return a tidy DataFrame with ``query_time``, ``log_prob``, and ``nll``.""" + pandas_module = _require_pandas_module() + rows = [ + {"query_time": query_time, "log_prob": score, "nll": negated} + for query_time, score, negated in zip( + self.query_times, + self.log_prob.values, + self.log_prob.nll_values, + strict=True, + ) + ] + return pandas_module.DataFrame.from_records(rows) + + def to_pandas_wide(self) -> Any: + """Return the one-row request summary over the scored horizons. + + The other result classes widen the column axis, which a log density does + not have. The wide view is therefore the request-level summary the + service reports, and the per-horizon scores stay in the tidy view. + """ + pandas_module = _require_pandas_module() + return pandas_module.DataFrame.from_records( + [ + { + "total": self.log_prob.total, + "mean": self.log_prob.mean, + "nll_total": self.log_prob.nll_total, + "nll_mean": self.log_prob.nll_mean, + } + ] + ) + + def build_forecast_payload( *, model_version: str, @@ -1384,6 +1645,7 @@ def build_forecast_payload( schema_version: str = SCHEMA_VERSION, query_mode: QueryMode = "forecast", condition: ConditionBlock | None = None, + query_rows: Sequence[Mapping[str, Any]] | None = None, ) -> dict[str, Any]: """Build a validated JSON-compatible forecast request payload.""" return ForecastRequest( @@ -1401,6 +1663,7 @@ def build_forecast_payload( quantiles=quantiles, seed=seed, condition=condition, + query_rows=query_rows, ).to_payload() @@ -1630,18 +1893,27 @@ def _validate_history_declared_columns( raise ValueError(f"history_rows are missing time_column {schema.time_column!r}") -def _serialize_history_rows( - history_rows: Sequence[Mapping[str, Any]], +def _serialize_rows( + rows: Sequence[Mapping[str, Any]], + *, + field: str, ) -> list[dict[str, Any]]: return [ { - key: _to_json_compatible(value, field=f"history_rows[{index}].{key}") - for key, value in history_row.items() + key: _to_json_compatible(value, field=f"{field}[{index}].{key}") + for key, value in row.items() } - for index, history_row in enumerate(history_rows) + for index, row in enumerate(rows) ] +def _require_scored_value(value: Any, *, field: str) -> None: + if value is None: + raise ValueError(f"{field} must carry an observed value to score") + if isinstance(value, float) and not math.isfinite(value): + raise ValueError(f"{field} must be finite") + + def _serialize_query_times( query_times: Sequence[Any], *, @@ -1979,6 +2251,38 @@ def _require_output_float(value: Any, *, field: str) -> float: return parsed +def _require_float_vector( + value: Any, + *, + field: str, + expected_length: int | None = None, +) -> tuple[float, ...]: + values = _require_sequence(value, field=field) + _validate_length( + actual_length=len(values), + expected_length=expected_length, + field=field, + ) + return tuple( + _require_output_float(item, field=f"{field}[{index}]") + for index, item in enumerate(values) + ) + + +def _require_derived_score(value: float, *, expected: float, field: str) -> float: + if not math.isclose( + value, + expected, + rel_tol=DERIVED_SCORE_TOLERANCE, + abs_tol=DERIVED_SCORE_TOLERANCE, + ): + raise ValueError( + f"{field} disagrees with outputs.log_prob.values: " + f"expected {expected}, got {value}" + ) + return value + + def _require_float_matrix( value: Any, *, diff --git a/tests/fixtures/condition_log_prob_request.json b/tests/fixtures/condition_log_prob_request.json new file mode 100644 index 0000000..2407223 --- /dev/null +++ b/tests/fixtures/condition_log_prob_request.json @@ -0,0 +1,51 @@ +{ + "schema_version": "v3", + "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", + "query_mode": "condition", + "return_mode": "log_prob", + "time_index_mode": "ordinal", + "columns": [ + { + "name": "driver", + "modality": "numeric" + }, + { + "name": "target", + "modality": "numeric", + "role": "target" + } + ], + "history_rows": [ + { + "driver": 1.0, + "target": 10.0 + }, + { + "driver": 1.1, + "target": 11.0 + } + ], + "query_times": [2, 3], + "requested_columns": ["target"], + "condition": { + "query_time_index": 1, + "conditions": [ + { + "column": "driver", + "kind": "equality", + "value": 1.5 + } + ] + }, + "query_rows": [ + { + "driver": 1.2, + "target": 12.0 + }, + { + "driver": 1.5, + "target": 13.0 + } + ], + "seed": 7 +} diff --git a/tests/fixtures/condition_log_prob_response.json b/tests/fixtures/condition_log_prob_response.json new file mode 100644 index 0000000..21d9f62 --- /dev/null +++ b/tests/fixtures/condition_log_prob_response.json @@ -0,0 +1,36 @@ +{ + "schema_version": "v3", + "image_version": "0.3.0", + "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", + "checkpoint_version": "sdk-test", + "head": "gmm", + "query_mode": "condition", + "return_mode": "log_prob", + "outputs": { + "query_times": [3], + "requested_columns": ["target"], + "mean": null, + "samples": null, + "quantiles": null, + "log_prob": { + "values": [-1.75], + "nll_values": [1.75], + "total": -1.75, + "mean": -1.75, + "nll_total": 1.75, + "nll_mean": 1.75 + } + }, + "plausibility": { + "equality_log_density": -1.27, + "region_log_probability": null + }, + "diagnostics": { + "history_rows": 2, + "horizon_count": 1, + "seed": 7, + "condition_draws": 500, + "interval_estimator": null + }, + "errors": [] +} diff --git a/tests/fixtures/forecast_log_prob_request.json b/tests/fixtures/forecast_log_prob_request.json new file mode 100644 index 0000000..f843ce4 --- /dev/null +++ b/tests/fixtures/forecast_log_prob_request.json @@ -0,0 +1,41 @@ +{ + "schema_version": "v3", + "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", + "query_mode": "forecast", + "return_mode": "log_prob", + "time_index_mode": "ordinal", + "columns": [ + { + "name": "driver", + "modality": "numeric" + }, + { + "name": "target", + "modality": "numeric", + "role": "target" + } + ], + "history_rows": [ + { + "driver": 1.0, + "target": 10.0 + }, + { + "driver": 1.1, + "target": 11.0 + } + ], + "query_times": [2, 3], + "requested_columns": ["driver", "target"], + "query_rows": [ + { + "driver": 1.2, + "target": 12.0 + }, + { + "driver": 1.3, + "target": 13.0 + } + ], + "seed": 7 +} diff --git a/tests/fixtures/forecast_log_prob_response.json b/tests/fixtures/forecast_log_prob_response.json new file mode 100644 index 0000000..f6ad94a --- /dev/null +++ b/tests/fixtures/forecast_log_prob_response.json @@ -0,0 +1,33 @@ +{ + "schema_version": "v3", + "image_version": "0.3.0", + "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", + "checkpoint_version": "sdk-test", + "head": "gmm", + "query_mode": "forecast", + "return_mode": "log_prob", + "outputs": { + "query_times": [2, 3], + "requested_columns": ["driver", "target"], + "mean": null, + "samples": null, + "quantiles": null, + "log_prob": { + "values": [-2.5, -3.25], + "nll_values": [2.5, 3.25], + "total": -5.75, + "mean": -2.875, + "nll_total": 5.75, + "nll_mean": 2.875 + } + }, + "plausibility": null, + "diagnostics": { + "history_rows": 2, + "horizon_count": 2, + "seed": 7, + "condition_draws": null, + "interval_estimator": null + }, + "errors": [] +} diff --git a/tests/test_fixture_compatibility.py b/tests/test_fixture_compatibility.py index e699ef8..f7a79cf 100644 --- a/tests/test_fixture_compatibility.py +++ b/tests/test_fixture_compatibility.py @@ -26,6 +26,7 @@ HealthMetadata, SCHEMA_VERSION, JointFMServiceError, + LogProbResult, MeanForecastResult, QuantileForecastResult, SampleForecastResult, @@ -78,6 +79,18 @@ def test_health_fixture_matches_current_service_contract( QuantileForecastResult, "quantiles", ), + ( + "forecast_log_prob_request", + "forecast_log_prob_response", + LogProbResult, + "log_prob", + ), + ( + "condition_log_prob_request", + "condition_log_prob_response", + LogProbResult, + "log_prob", + ), ], ) def test_checked_in_fixture_payloads_parse_as_forecast_results( diff --git a/tests/test_log_prob_mode.py b/tests/test_log_prob_mode.py new file mode 100644 index 0000000..cdb6e01 --- /dev/null +++ b/tests/test_log_prob_mode.py @@ -0,0 +1,363 @@ +# Copyright 2026 DataRobot, Inc. and its affiliates. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Tests for the ``log_prob`` return mode on the client side. + +Scoring inverts the direction of every other return mode: the caller supplies +the values and the deployment answers how plausible they are. That shows up in +three places, and the split below follows them. The request carries observed +rows nothing else carries, and a projection narrower than the scored joint asks +a different question than the caller believes they asked. The response carries +a block with no column axis whose summary fields only repeat what ``values`` +already says. Under a condition the two meet: a pinned column leaves the +read-out set but still has to be supplied, because the service refuses a row +contradicting the pin rather than scoring it. +""" + +from __future__ import annotations + +from collections.abc import Callable, Mapping +from typing import Any + +import pytest + +from jointfm_client import ( + ColumnSpec, + ConditionBlock, + ConditionPlausibility, + DataFrameSchema, + EqualityCondition, + ForecastRequest, + ForecastRequestMetadata, + ForecastResponse, + JointFMClient, + LogProbResult, + build_forecast_payload_from_dataframe, +) +from jointfm_client.contract import ReturnMode + +_MODEL_VERSION = "jointfm-inference:0.3.0+ckpt.sdk-test" +_HISTORY_ROWS = ( + {"driver": 1.0, "target": 10.0}, + {"driver": 1.1, "target": 11.0}, +) +_QUERY_ROWS = ( + {"driver": 1.2, "target": 12.0}, + {"driver": 1.3, "target": 13.0}, +) +_PINNED_QUERY_ROWS = ( + {"driver": 1.2, "target": 12.0}, + {"driver": 1.5, "target": 13.0}, +) +_PIN_DRIVER = ConditionBlock( + query_time_index=1, + conditions=(EqualityCondition(column="driver", value=1.5),), +) + + +def _schema() -> DataFrameSchema: + """Build one two-column ordinal schema for request-level tests.""" + return DataFrameSchema( + columns=( + ColumnSpec(name="driver", modality="numeric"), + ColumnSpec(name="target", modality="numeric", role="target"), + ), + time_index_mode="ordinal", + ) + + +def _request( + *, + return_mode: ReturnMode = "log_prob", + query_rows: Any = _QUERY_ROWS, + requested_columns: tuple[str, ...] | None = ("driver", "target"), + condition: ConditionBlock | None = None, +) -> ForecastRequest: + """Build one scoring request against the two-column schema.""" + return ForecastRequest( + metadata=ForecastRequestMetadata( + model_version=_MODEL_VERSION, + query_mode="forecast" if condition is None else "condition", + return_mode=return_mode, + ), + schema=_schema(), + history_rows=_HISTORY_ROWS, + query_times=(2, 3), + requested_columns=requested_columns, + condition=condition, + query_rows=query_rows, + seed=7, + ) + + +class _ScoringTransport: + """Fake JSON transport for a deployment that serves one canned score.""" + + def __init__( + self, + *, + health_payload: Mapping[str, Any], + predict_payload: Mapping[str, Any], + ) -> None: + """Remember the advertisement and the canned score to serve.""" + self.health_payload = health_payload + self.predict_payload = predict_payload + self.payloads: list[dict[str, Any]] = [] + + def get_json(self, url: str) -> Mapping[str, Any]: + """Serve the health advertisement on the local health route.""" + assert url == "http://127.0.0.1:8080/healthz" + return self.health_payload + + def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: + """Record the scoring payload and answer it.""" + assert url == "http://127.0.0.1:8080/predict" + self.payloads.append(dict(payload)) + return self.predict_payload + + +def _client(transport: _ScoringTransport) -> JointFMClient: + """Build a local-service client over the fake transport.""" + return JointFMClient( + health_url="http://127.0.0.1:8080/healthz", + predict_url="http://127.0.0.1:8080/predict", + transport=transport, + ) + + +def test_scoring_needs_the_values_it_scores() -> None: + """The mode and the observed rows are one request, not two independent options.""" + with pytest.raises(ValueError, match="go together"): + _request(query_rows=None) + + with pytest.raises(ValueError, match="go together"): + _request(return_mode="mean", requested_columns=("target",)) + + +def test_every_future_position_needs_its_observed_row() -> None: + """One score exists per query time, so a missing row has no answer to give.""" + with pytest.raises(ValueError, match="one row per query time"): + _request(query_rows=(_QUERY_ROWS[0],)) + + +@pytest.mark.parametrize( + ("row", "message"), + [ + ({"driver": 1.3}, "missing declared columns"), + ({"driver": 1.3, "target": None}, "must carry an observed value"), + ({"driver": 1.3, "target": float("nan")}, "must be finite"), + ], +) +def test_a_row_that_cannot_be_scored_is_refused_before_any_round_trip( + row: Mapping[str, Any], + message: str, +) -> None: + """The score is joint over every declared column, so each one needs a value.""" + with pytest.raises(ValueError, match=message): + _request(query_rows=(_QUERY_ROWS[0], row)) + + +def test_a_partial_projection_is_refused_because_it_scores_something_else() -> None: + """Dropping a column would score a narrower joint than the caller believes.""" + with pytest.raises(ValueError, match="every declared column in schema order"): + _request(requested_columns=("target",)) + + +def test_a_pinned_column_is_the_one_column_a_score_may_leave_out() -> None: + """Conditioning fixes the column, and the contract forbids reading it back.""" + pinned = _request( + requested_columns=("target",), + query_rows=_PINNED_QUERY_ROWS, + condition=_PIN_DRIVER, + ) + assert pinned.to_payload()["requested_columns"] == ["target"] + + default_projection = _request( + requested_columns=None, + query_rows=_PINNED_QUERY_ROWS, + condition=_PIN_DRIVER, + ) + assert "requested_columns" not in default_projection.to_payload() + + with pytest.raises(ValueError, match="lists pinned columns"): + _request(query_rows=_PINNED_QUERY_ROWS, condition=_PIN_DRIVER) + + +@pytest.mark.parametrize( + ("fixture_name", "condition", "requested_columns", "query_rows"), + [ + ("forecast_log_prob_request", None, ("driver", "target"), _QUERY_ROWS), + ( + "condition_log_prob_request", + _PIN_DRIVER, + ("target",), + _PINNED_QUERY_ROWS, + ), + ], +) +def test_the_payload_carries_the_rows_the_service_scores( + json_fixture_loader: Callable[[str], dict[str, Any]], + fixture_name: str, + condition: ConditionBlock | None, + requested_columns: tuple[str, ...], + query_rows: tuple[Mapping[str, Any], ...], +) -> None: + """What the request builds is what the checked-in service fixture contains.""" + payload = _request( + requested_columns=requested_columns, + condition=condition, + query_rows=query_rows, + ).to_payload() + + assert payload == json_fixture_loader(fixture_name) + + +def test_the_scored_response_parses_into_its_own_result( + json_fixture_loader: Callable[[str], dict[str, Any]], +) -> None: + """A score has a horizon axis and no column axis, and says so in its result.""" + result = ForecastResponse.from_payload( + json_fixture_loader("forecast_log_prob_response"), + request_payload=json_fixture_loader("forecast_log_prob_request"), + ) + + assert isinstance(result, LogProbResult) + assert result.log_prob.values == (-2.5, -3.25) + assert result.log_prob.nll_values == (2.5, 3.25) + assert result.log_prob.total == pytest.approx(-5.75) + assert result.log_prob.mean == pytest.approx(-2.875) + assert result.plausibility is None + assert result.outputs.log_prob is result.log_prob + + +def test_the_scored_conversions_keep_the_horizon_axis( + json_fixture_loader: Callable[[str], dict[str, Any]], +) -> None: + """Tidy carries one row per scored position; wide carries the request summary.""" + pytest.importorskip("pandas") + pytest.importorskip("numpy") + result = ForecastResponse.from_payload( + json_fixture_loader("forecast_log_prob_response"), + request_payload=json_fixture_loader("forecast_log_prob_request"), + ) + assert isinstance(result, LogProbResult) + + assert result.to_numpy().tolist() == [-2.5, -3.25] + tidy = result.to_pandas_tidy() + assert list(tidy.columns) == ["query_time", "log_prob", "nll"] + assert tidy["query_time"].tolist() == [2, 3] + wide = result.to_pandas_wide() + assert list(wide.columns) == ["total", "mean", "nll_total", "nll_mean"] + assert len(wide) == 1 + + +@pytest.mark.parametrize("field", ["total", "mean", "nll_total", "nll_mean"]) +def test_a_summary_that_disagrees_with_its_scores_is_refused( + json_fixture_loader: Callable[[str], dict[str, Any]], + field: str, +) -> None: + """The summaries are redundant, so a payload that disagrees with itself is broken.""" + payload = json_fixture_loader("forecast_log_prob_response") + payload["outputs"]["log_prob"][field] += 1.0 + + with pytest.raises(ValueError, match=f"outputs.log_prob.{field} disagrees"): + ForecastResponse.from_payload(payload) + + +def test_a_scored_response_is_never_read_as_another_mode( + json_fixture_loader: Callable[[str], dict[str, Any]], +) -> None: + """A mode whose block is absent must fail on its own field, not on another's.""" + payload = json_fixture_loader("forecast_mean_response") + payload["return_mode"] = "log_prob" + + with pytest.raises(ValueError, match="outputs.log_prob"): + ForecastResponse.from_payload(payload) + + +def test_the_client_scores_observed_rows_and_types_the_answer( + json_fixture_loader: Callable[[str], dict[str, Any]], +) -> None: + """The typed helper puts the rows on the wire and types what comes back.""" + transport = _ScoringTransport( + health_payload=json_fixture_loader("health_metadata"), + predict_payload=json_fixture_loader("forecast_log_prob_response"), + ) + + result = _client(transport).forecast_log_prob( + list(_HISTORY_ROWS), + schema=_schema(), + query_times=[2, 3], + query_rows=list(_QUERY_ROWS), + requested_columns=["driver", "target"], + model_version=_MODEL_VERSION, + seed=7, + ) + + assert isinstance(result, LogProbResult) + assert result.log_prob.values == (-2.5, -3.25) + sent = transport.payloads[0] + assert sent["return_mode"] == "log_prob" + assert sent["query_rows"] == [dict(row) for row in _QUERY_ROWS] + + +def test_the_client_scores_under_a_condition( + json_fixture_loader: Callable[[str], dict[str, Any]], +) -> None: + """A conditioned score answers at the conditioned position and reports its pin.""" + health_payload = json_fixture_loader("health_metadata") + health_payload["supported_query_modes"] = ["forecast", "condition"] + health_payload["supported_condition_kinds"] = ["equality", "interval"] + transport = _ScoringTransport( + health_payload=health_payload, + predict_payload=json_fixture_loader("condition_log_prob_response"), + ) + + result = _client(transport).forecast_log_prob( + list(_HISTORY_ROWS), + schema=_schema(), + query_times=[2, 3], + query_rows=list(_PINNED_QUERY_ROWS), + requested_columns=["target"], + model_version=_MODEL_VERSION, + seed=7, + condition=_PIN_DRIVER, + ) + + assert isinstance(result, LogProbResult) + assert result.query_times == (3,) + assert result.log_prob.values == (-1.75,) + assert result.plausibility == ConditionPlausibility(equality_log_density=-1.27) + sent = transport.payloads[0] + assert sent["query_mode"] == "condition" + assert sent["query_rows"] == [dict(row) for row in _PINNED_QUERY_ROWS] + + +def test_the_dataframe_adapter_encodes_observed_rows_like_history() -> None: + """A pandas caller gets both arrays through the same column encoding.""" + pandas = pytest.importorskip("pandas") + + payload = build_forecast_payload_from_dataframe( + pandas.DataFrame(list(_HISTORY_ROWS)), + model_version=_MODEL_VERSION, + time_index_mode="ordinal", + query_times=[2, 3], + target_columns=["target"], + requested_columns=["driver", "target"], + return_mode="log_prob", + query_rows=pandas.DataFrame(list(_QUERY_ROWS)), + ) + + assert payload["return_mode"] == "log_prob" + assert payload["query_rows"] == [dict(row) for row in _QUERY_ROWS] From fb6c0186fe5e6fd2628ecfc7f008e1c201cffe1b Mon Sep 17 00:00:00 2001 From: Stefan Hackmann Date: Wed, 23 Sep 2026 08:40:14 +0000 Subject: [PATCH 2/8] docs: Lead agent responses with the outcome and keep the evidence below --- AGENTS.md | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/AGENTS.md b/AGENTS.md index 823357d..faeeb9d 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -1 +1 @@ -Agents must ground their thoughts in facts, not assumptions: before planning, claiming, or editing, read the relevant material — source, configs, data, docs, test output — and only act on beliefs backed by something just read or run. Agents must run ad-hoc Python and CLI commands through `task run -- ` (which wraps `uv run` with the project's canonical parameters), never bare `python` or hand-written `uv run` invocations — this keeps every invocation from accidentally re-resolving or mutating the `.venv` and `uv.lock`. Agents must explain what each command will do and why it is being run before running it. Before asking the user a question, agents must first explain the corresponding context and terminology — what the question concerns, why it arises, and what any project-specific terms mean — so the user can answer without digging through the code themselves. Agents must run `task pre-commit` and fix all reported issues before reporting success to the user. After `task pre-commit` succeeds, show the diff and explain why each change is necessary before reporting success to the user. When what an agent reports is a set of items measured or classified along the same few dimensions — files and their status, config keys and their values, metrics before and after, options and their trade-offs — it must be presented as a Markdown table instead of prose or a bullet list, with short cells and only the columns that carry information; prose stays the right form for a single item, a narrative explanation, or an ordered procedure, because a one-row table, or one whose cells are sentences, is harder to read than the paragraph it replaced. Agents must never perform destructive or state-changing git operations unless the user explicitly instructs the agent to run that specific operation — no `git push --force`, no `git reset --hard`, no `git stash` (which hides uncommitted work), no branch/tag deletion, no history rewrites (`rebase`, `commit --amend` on published commits, `filter-branch`), no `git clean -fdx`, no `--no-verify` to bypass hooks, and no discarding of uncommitted work. Read-only inspection commands (`git status`, `git diff`, `git log`, `git show`) are always allowed. If a task seems to require a state-changing git operation and the user has not explicitly asked for it, stop and ask the user to run it. Agents must never mention temporary planning identifiers — phase names or numbers such as "Phase 3g", milestone, sprint, or ticket codes — in docstrings, comments, help or description strings, error messages, or test docstrings; those labels are deleted when the plan is retired and leave readers with a dangling reference nobody can decode, so describe the concept by its lasting behavior or stable configuration key instead (phase names belong only in roadmap and planning docs, and linking to such a doc by its actual filename is fine). Agents must not mention or recommend key rotation (rotating API keys, tokens, or other credentials) — the user takes care of key rotation themselves. Agents must keep comments in configuration files value-independent: a comment attached to a configuration parameter — in `config.yaml`/`config.sample.yaml`, a `Taskfile`, `pyproject.toml`, or any other settings file — may only state what remains true whatever the value is (what the parameter controls, what its options mean, which invariant or trade-off it participates in), and must never restate the current value, derive from it, or do arithmetic on it (`# 64 of 576 rows`, `# 2x the 288-step horizon`, `# keeps 90% of draws at full context`); configuration is retuned constantly, so such a comment turns false the moment somebody edits the value on the line below it and nothing catches the drift, so express a needed derivation as a rule over the parameter or record concrete numbers in a dated notes doc where they read as a historical observation. ripgrep (`rg`) is available (installed by `task setup` via `task install:ripgrep`); prefer it for fast code and text search. +Agents must ground their thoughts in facts, not assumptions: before planning, claiming, or editing, read the relevant material — source, configs, data, docs, test output — and only act on beliefs backed by something just read or run. Agents must run ad-hoc Python and CLI commands through `task run -- ` (which wraps `uv run` with the project's canonical parameters), never bare `python` or hand-written `uv run` invocations — this keeps every invocation from accidentally re-resolving or mutating the `.venv` and `uv.lock`. Agents must explain what each command will do and why it is being run before running it. Before asking the user a question, agents must first explain the corresponding context and terminology — what the question concerns, why it arises, and what any project-specific terms mean — so the user can answer without digging through the code themselves. Agents must run `task pre-commit` and fix all reported issues before reporting success to the user. After `task pre-commit` succeeds, show the diff and explain why each change is necessary before reporting success to the user. Agents must lead a response with its outcome — what now works, what the answer is, and anything the user must decide or act on — within the first few lines, so a reader who stops there misses nothing that changes what they do next; the evidence and the mechanism follow underneath, ordered by what the user cannot reconstruct alone, and what the user can already see is left out: no restatement of the request, no narration of the steps taken or the files visited, no walk through a diff that is shown anyway, and no closing summary that repeats the opening. That focus governs only what is *not* written — it never removes a `module.py:120` citation, a measured number, a side-effect or migration impact, an unrelated defect the agent noticed, or a question the user must answer; when focus and justification collide, the prose goes and the evidence stays. When what an agent reports is a set of items measured or classified along the same few dimensions — files and their status, config keys and their values, metrics before and after, options and their trade-offs — it must be presented as a Markdown table instead of prose or a bullet list, with short cells and only the columns that carry information; prose stays the right form for a single item, a narrative explanation, or an ordered procedure, because a one-row table, or one whose cells are sentences, is harder to read than the paragraph it replaced. Agents must never perform destructive or state-changing git operations unless the user explicitly instructs the agent to run that specific operation — no `git push --force`, no `git reset --hard`, no `git stash` (which hides uncommitted work), no branch/tag deletion, no history rewrites (`rebase`, `commit --amend` on published commits, `filter-branch`), no `git clean -fdx`, no `--no-verify` to bypass hooks, and no discarding of uncommitted work. Read-only inspection commands (`git status`, `git diff`, `git log`, `git show`) are always allowed. If a task seems to require a state-changing git operation and the user has not explicitly asked for it, stop and ask the user to run it. Agents must never mention temporary planning identifiers — phase names or numbers such as "Phase 3g", milestone, sprint, or ticket codes — in docstrings, comments, help or description strings, error messages, or test docstrings; those labels are deleted when the plan is retired and leave readers with a dangling reference nobody can decode, so describe the concept by its lasting behavior or stable configuration key instead (phase names belong only in roadmap and planning docs, and linking to such a doc by its actual filename is fine). Agents must not mention or recommend key rotation (rotating API keys, tokens, or other credentials) — the user takes care of key rotation themselves. Agents must keep comments in configuration files value-independent: a comment attached to a configuration parameter — in `config.yaml`/`config.sample.yaml`, a `Taskfile`, `pyproject.toml`, or any other settings file — may only state what remains true whatever the value is (what the parameter controls, what its options mean, which invariant or trade-off it participates in), and must never restate the current value, derive from it, or do arithmetic on it (`# 64 of 576 rows`, `# 2x the 288-step horizon`, `# keeps 90% of draws at full context`); configuration is retuned constantly, so such a comment turns false the moment somebody edits the value on the line below it and nothing catches the drift, so express a needed derivation as a rule over the parameter or record concrete numbers in a dated notes doc where they read as a historical observation. ripgrep (`rg`) is available (installed by `task setup` via `task install:ripgrep`); prefer it for fast code and text search. From 326452b80267bf4fd07c33b8a46549dad5b72856 Mon Sep 17 00:00:00 2001 From: Stefan Hackmann Date: Wed, 23 Sep 2026 10:31:45 +0000 Subject: [PATCH 3/8] fix: Answer every declared column by default and move the envelope to v4 --- README.md | 20 +++---- config.sample.yaml | 2 +- docs/api-reference.md | 16 ++--- src/jointfm_client/contract.py | 42 +++++--------- .../fixtures/condition_interval_response.json | 2 +- .../fixtures/condition_log_prob_request.json | 4 +- .../fixtures/condition_log_prob_response.json | 4 +- tests/fixtures/condition_mean_request.json | 2 +- tests/fixtures/condition_mean_response.json | 2 +- tests/fixtures/forecast_log_prob_request.json | 2 +- .../fixtures/forecast_log_prob_response.json | 2 +- tests/fixtures/forecast_mean_request.json | 2 +- tests/fixtures/forecast_mean_response.json | 2 +- .../fixtures/forecast_quantiles_request.json | 2 +- .../fixtures/forecast_quantiles_response.json | 2 +- tests/fixtures/forecast_samples_request.json | 2 +- tests/fixtures/forecast_samples_response.json | 2 +- tests/fixtures/health_metadata.json | 2 +- .../input_size_exceeded_response.json | 2 +- .../model_version_mismatch_response.json | 2 +- .../schema_version_mismatch_response.json | 4 +- tests/fixtures/validation_error_response.json | 2 +- tests/test_cli.py | 6 +- tests/test_condition_mode.py | 46 +++++++++++---- tests/test_configuration.py | 6 +- tests/test_contract.py | 24 ++++---- tests/test_contract_models.py | 12 ++-- tests/test_feature_importance.py | 2 +- tests/test_log_prob_mode.py | 35 ++++++----- tests/test_pool.py | 18 +++--- tests/test_settings.py | 12 ++-- tests/test_surfaces.py | 4 +- tests/test_transport.py | 58 +++++++++---------- 33 files changed, 181 insertions(+), 164 deletions(-) diff --git a/README.md b/README.md index 50eab26..2b1e7f0 100644 --- a/README.md +++ b/README.md @@ -10,7 +10,7 @@ The SDK targets the DataRobot-hosted unstructured prediction route and the same - Import namespace: `jointfm_client` - Supported Python: `>=3.11` - Current SDK package version: `0.8.0` -- Current JointFM service schema: `schema_version="v3"` +- Current JointFM service schema: `schema_version="v4"` The public API shape is a synchronous low-level `JointFMClient` with `health()`, `health_instances()`, and `predict(payload)` methods plus high-level `forecast(...)`, `forecast_mean(...)`, `forecast_samples(...)`, and `forecast_quantiles(...)` helpers. The SDK is not a proxy service; callers use it as a local Python library that talks to the hosted or local JointFM endpoint. @@ -50,7 +50,7 @@ Example deployment configuration: deployment: datarobot_endpoint: https://app.datarobot.com/api/v2 datarobot_api_token: - schema_version: v3 + schema_version: v4 deployment_id: # Optional model-version pin; the SDK discovers it from /healthz when unset: # model_version: jointfm-inference:0.3.0+ckpt.fin-2026-05-22 @@ -68,7 +68,7 @@ Equivalent `.env` deployment configuration: ```dotenv DATAROBOT_ENDPOINT=https://app.datarobot.com/api/v2 DATAROBOT_API_TOKEN= -JOINTFM_SCHEMA_VERSION=v3 +JOINTFM_SCHEMA_VERSION=v4 JOINTFM_DEPLOYMENT_ID= # Optional drift-detection pin; the SDK discovers the model version from /healthz when unset: # JOINTFM_MODEL_VERSION=jointfm-inference:0.3.0+ckpt.fin-2026-05-22 @@ -78,7 +78,7 @@ Equivalent local REST configuration for a service started from the `joint` repos ```dotenv JOINTFM_LOCAL_BASE_URL=http://127.0.0.1:8080 -JOINTFM_SCHEMA_VERSION=v3 +JOINTFM_SCHEMA_VERSION=v4 # Optional drift-detection pin; the SDK discovers the model version from /healthz when unset: # JOINTFM_MODEL_VERSION=jointfm-inference:0.3.0+ckpt.fin_i504_o63_f0_t10_h16l16_mam7_af_t3r1_cnn_k3l4_hpst_h16l2_studentt_m4cr2df8skew ``` @@ -167,7 +167,7 @@ The bootstrap helper resolves the nearest src-layout Python project root, switch The current forecast request contract is: -- `schema_version`: exactly `"v3"`, configured as `JOINTFM_SCHEMA_VERSION` for `from_env()` clients +- `schema_version`: exactly `"v4"`, configured as `JOINTFM_SCHEMA_VERSION` for `from_env()` clients - `model_version`: exact model version advertised by `/healthz` or otherwise selected by the caller. Optional for `from_env()` clients: when `JOINTFM_MODEL_VERSION` is unset the SDK reads it from `/healthz` on first use; when set it acts as a drift-detection pin - `query_mode`: `"forecast"` for the unconditional forecast, or `"condition"` for a conditional query at one future position; the high-level helpers set it from whether a `condition` block was passed - `return_mode`: one of `"mean"`, `"samples"`, or `"quantiles"` @@ -184,10 +184,10 @@ The `condition` query mode asks for the model's joint distribution at one future A column carries at most one condition, of either kind: -- `EqualityCondition(column, value)` pins the column to a finite value. The pinned column leaves the read-out set, so it must not appear in `requested_columns`. +- `EqualityCondition(column, value)` pins the column to a finite value. It stays nameable in `requested_columns` and reads back the value the request supplied, so a scenario answer lines up column for column with an unconditioned one. - `IntervalCondition(column, lower=None, upper=None)` confines the column to a range; `None` leaves that side open, and at least one side must be bounded. The column stays readable, and what comes back is its distribution inside the range. -Every column without a condition is a read-out column, and at least one must remain. Pass the block to `forecast(...)`, `forecast_mean(...)`, `forecast_samples(...)`, or `forecast_quantiles(...)`: +Every column without a condition is a read-out column, and at least one must remain. `requested_columns` chooses what the response carries independently of that, defaulting to every declared column in declared order. Pass the block to `forecast(...)`, `forecast_mean(...)`, `forecast_samples(...)`, or `forecast_quantiles(...)`: ```python from jointfm_client import ConditionBlock, EqualityCondition, IntervalCondition @@ -227,7 +227,7 @@ Successful forecast responses preserve `schema_version`, `image_version`, `model ```json { - "schema_version": "v3", + "schema_version": "v4", "errors": [ { "code": "VALIDATION_ERROR", @@ -242,7 +242,7 @@ Known error codes are `VALIDATION_ERROR`, `UNSUPPORTED_HEAD_QUERY_COMBINATION`, ## Compatibility Policy -The SDK supports only `schema_version="v3"`. `validate_service_metadata()` checks `/healthz` metadata and raises typed compatibility errors before prediction if the service advertises a different schema, an unexpected model version, mode capabilities outside the recorded service contract, or an unsupported `decoding_strategy`. Return modes and time-index modes must match the SDK's lists exactly. Query modes and condition kinds are derived by the service from the mounted head, so a deployment may advertise fewer of them than the SDK knows; it must advertise at least one query mode and nothing the SDK does not know. +The SDK supports only `schema_version="v4"`. `validate_service_metadata()` checks `/healthz` metadata and raises typed compatibility errors before prediction if the service advertises a different schema, an unexpected model version, mode capabilities outside the recorded service contract, or an unsupported `decoding_strategy`. Return modes and time-index modes must match the SDK's lists exactly. Query modes and condition kinds are derived by the service from the mounted head, so a deployment may advertise fewer of them than the SDK knows; it must advertise at least one query mode and nothing the SDK does not know. Callers should pass an expected `model_version` when they already know which deployment artifact they intend to use. A mismatch is treated as a hard compatibility error rather than silently downgrading, guessing, or retrying another model. @@ -284,7 +284,7 @@ Create `.env` from `.env.sample` or set the same values in your shell. A hosted ```dotenv DATAROBOT_ENDPOINT=https://app.datarobot.com/api/v2 DATAROBOT_API_TOKEN= -JOINTFM_SCHEMA_VERSION=v3 +JOINTFM_SCHEMA_VERSION=v4 JOINTFM_DEPLOYMENT_ID= # Or: JOINTFM_DEPLOYMENT_IDS=chevron-id,research-id # Optional drift-detection pin; the SDK discovers the model version from /healthz when unset: diff --git a/config.sample.yaml b/config.sample.yaml index 768f4c5..a49d28e 100644 --- a/config.sample.yaml +++ b/config.sample.yaml @@ -52,7 +52,7 @@ transport: - X-DataRobot-Execution-ID user_agent_header: User-Agent forecast: - schema_version: v3 + schema_version: v4 query_mode: forecast return_mode: mean time_index_mode: ordinal diff --git a/docs/api-reference.md b/docs/api-reference.md index e8de616..efc7f82 100644 --- a/docs/api-reference.md +++ b/docs/api-reference.md @@ -18,7 +18,7 @@ This reference covers the supported public Python surface exported by `jointfm_c | `DataFrameSchema` | Describes tabular history layout. Fields are `columns`, `time_index_mode`, `time_column`, `time_scale_seconds`, `use_local_normalized_time`, `calendar_id`, and `timezone`. | | `ForecastRequestMetadata` | Holds `schema_version`, `model_version`, `query_mode` (`forecast` or `condition`), and `return_mode` for one forecast request. | | `ForecastRequest` | Validated request object that combines metadata, schema, history rows, query times, requested columns, sample or quantile controls, `seed`, an optional `condition` block, and the optional `query_rows` a scored request supplies, then emits a JSON-compatible payload with `to_payload()`. `query_mode="condition"` and a `condition` block must appear together, as must `return_mode="log_prob"` and `query_rows`, and both are validated against the request's own schema and `query_times` before any round trip. | -| `EqualityCondition` | One column of the conditioned position pinned to a finite `value`. The pinned column leaves the read-out set, so it must not appear in `requested_columns`. | +| `EqualityCondition` | One column of the conditioned position pinned to a finite `value`. The column may be named in `requested_columns` like any other, and reads back the pinned value the request supplied. | | `IntervalCondition` | One column of the conditioned position confined to `[lower, upper]`; `None` leaves a side open, at least one side must be bounded, and `lower < upper`. The column stays readable and the response describes its distribution inside the range. | | `ConditionBlock` | Every condition of one request: `query_time_index` into the request's `query_times` and a sequence of `EqualityCondition` or `IntervalCondition` values, at most one per column. `kinds`, `pinned_columns`, and `conditioned_columns` expose what the capability gate and the request validation need. | | `ConditionPlausibility` | What the model thinks of the conditions it was given: `equality_log_density` of the pinned values and `region_log_probability` of the interval region, each `None` when the request carried no condition of that kind. Reported by the service and never refused on. | @@ -80,7 +80,7 @@ All SDK-specific exceptions inherit from `JointFMError`. | `JointFMHTTPStatusError` | The service returns an HTTP error status. | | `JointFMServiceError` | A response body contains non-empty JointFM `errors`, including the case where HTTP status unexpectedly succeeded. | | `JointFMCompatibilityError` | Base class for fail-fast service compatibility failures. | -| `UnsupportedSchemaVersionError` | The service or response advertises a schema version other than `v3`. | +| `UnsupportedSchemaVersionError` | The service or response advertises a schema version other than `v4`. | | `UnsupportedModelVersionError` | The service or response model version differs from the configured or requested version. | | `UnsupportedServiceContractError` | The service-health payload advertises mode capabilities or a `decoding_strategy` outside the recorded service contract, or a condition request targets a deployment that does not advertise the `condition` mode or one of the block's condition kinds. | @@ -126,7 +126,7 @@ All SDK-specific exceptions inherit from `JointFMError`. | --- | --- | --- | | `DATAROBOT_ENDPOINT` | Hosted calls | HTTPS DataRobot API v2 endpoint, normalized without a trailing slash and required to end in `/api/v2`. | | `DATAROBOT_API_TOKEN` | Hosted calls | Non-empty, whitespace-free API token used in the hosted bearer authorization header. | -| `JOINTFM_SCHEMA_VERSION` | Hosted calls | Request schema pin. The SDK supports only `v3`. | +| `JOINTFM_SCHEMA_VERSION` | Hosted calls | Request schema pin. The SDK supports only `v4`. | | `JOINTFM_MODEL_VERSION` | Hosted calls | Exact JointFM deployment model version expected from the service-health payload and prediction responses. | | `JOINTFM_DEPLOYMENT_ID` | One selector | Deployment ID used to build hosted health and prediction URLs. | | `JOINTFM_DEPLOYMENT_IDS` | One selector | Comma-separated hosted deployment IDs for round-robin load balancing (at least two unique IDs). Mutually exclusive with other selectors. Peers must share `model_version` and `checkpoint_version`. `health()` uses the minimum reachable `max_sample_count` as the sample-batch cap; `health_instances()` sums reachable caps for overall parallel capacity and reports topology. | @@ -157,7 +157,7 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | Field | Required | Description | | --- | --- | --- | | `request_type` | Optional | One of `"predict"` (default) or `"health"`. Forecast requests omit this field or set it to `"predict"`. | -| `schema_version` | Yes | Must be `"v3"`. | +| `schema_version` | Yes | Must be `"v4"`. | | `model_version` | Yes | Exact deployed model version expected by the caller. | | `query_mode` | Yes | `"forecast"` for the unconditional forecast or `"condition"` for a conditional query at one future position. | | `return_mode` | Yes | One of `"mean"`, `"samples"`, `"quantiles"`, or `"log_prob"`, covered by the `forecast_mean`, `forecast_samples`, `forecast_quantiles`, and `forecast_log_prob` helpers. | @@ -166,12 +166,12 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | `history_rows` | Yes | Non-empty array of history row objects in the declared schema. | | `query_times` | Yes | Non-empty future forecast horizon values. Absolute datetimes are encoded timezone-stably. | | `time_column` | For absolute datetime, optional otherwise | Name of the history time column. It must not duplicate a modeled column name. | -| `requested_columns` | Optional | Output column names or integer indices. Duplicates are rejected. Defaults to all modeled columns, minus the columns an equality condition pinned. Under `return_mode="log_prob"` an explicit list must name every readable column in schema order, because the score is joint over them. | +| `requested_columns` | Optional | Output column names or integer indices, answered in the order given. Duplicates are rejected. Any declared column may be named, whatever its `role` and whatever a condition says about it. Defaults to every declared column in declared order, so a condition response and a forecast response over the same schema line up column for column. Under `return_mode="log_prob"` an explicit list must name every declared column in declared order, because the score is joint over them. | | `query_rows` | `log_prob` mode | Observed rows to score, one per entry of `query_times`, each carrying a value for every declared column. Required by `return_mode="log_prob"` and rejected for every other mode. A row that contradicts a condition — a pinned column given another value, or a bounded column outside its range — is refused by the service instead of scored. | | `n_samples` | Samples and quantiles controls | Positive sample count when sampling controls are needed. Oversized sample forecasts are batched automatically against the cap advertised in health metadata. | | `quantiles` | Quantiles mode | Quantile levels in `(0, 1)`, required for `return_mode="quantiles"`. | | `seed` | Optional | Integer random seed for reproducible stochastic outputs. | -| `condition` | With `query_mode="condition"` | Object with `query_time_index` (index into `query_times`) and `conditions`, a list of `{"column", "kind": "equality", "value"}` or `{"column", "kind": "interval", "lower", "upper"}` entries with `null` for an open bound. At most one condition per column, at least one column left unconditioned, and no pinned column in `requested_columns`. Forbidden with any other query mode. | +| `condition` | With `query_mode="condition"` | Object with `query_time_index` (index into `query_times`) and `conditions`, a list of `{"column", "kind": "equality", "value"}` or `{"column", "kind": "interval", "lower", "upper"}` entries with `null` for an open bound. At most one condition per column and at least one column left unconditioned. Forbidden with any other query mode. | | `time_scale_seconds` | Optional | Positive scale for continuous time indexes. | | `use_local_normalized_time` | Optional | Whether the service should use local normalized time features. | | `calendar_id` | Optional | Calendar identifier, defaulting to `pandas-default`. | @@ -200,7 +200,7 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | Field | Description | | --- | --- | -| `schema_version` | Response schema, expected to be `"v3"`. | +| `schema_version` | Response schema, expected to be `"v4"`. | | `image_version` | Service image version that produced the response. | | `model_version` | Model version that produced the response. | | `checkpoint_version` | Checkpoint version that produced the response. | @@ -226,7 +226,7 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | Field | Description | | --- | --- | | `status` | Service status string. | -| `schema_version` | Advertised schema version. The SDK requires `v3`. | +| `schema_version` | Advertised schema version. The SDK requires `v4`. | | `image_version` | Running service image version. | | `model_version` | Running model version. | | `checkpoint_version` | Loaded checkpoint version. | diff --git a/src/jointfm_client/contract.py b/src/jointfm_client/contract.py index 208a1d9..2b028b4 100644 --- a/src/jointfm_client/contract.py +++ b/src/jointfm_client/contract.py @@ -37,7 +37,7 @@ # pyproject.toml the way a hand-maintained literal here did. PACKAGE_VERSION: Final = importlib.metadata.version(DISTRIBUTION_NAME) -SCHEMA_VERSION: Final = "v3" +SCHEMA_VERSION: Final = "v4" # The service sums and averages its log densities in the model's own precision, # so its summaries differ from a recomputation in the last few digits. DERIVED_SCORE_TOLERANCE: Final = 1e-6 @@ -317,8 +317,9 @@ class EqualityCondition: An equality condition is an event of probability zero under a continuous head, so the deployment answers it analytically rather than by filtering - draws. It fixes the column's value, which is why that column leaves the - read-out projection. + draws. It fixes the column's value, so a projection that names the column + reads that value back rather than a model output — which is what lets a + scenario answer line up column for column with an unconditioned one. """ column: str @@ -423,7 +424,7 @@ def __post_init__(self) -> None: @property def pinned_columns(self) -> tuple[str, ...]: - """Columns an equality condition fixes, which leave the read-out set.""" + """Columns an equality condition fixes, whose answer is the request's own value.""" return tuple( condition.column for condition in self.conditions @@ -581,29 +582,17 @@ def _validate_condition(self) -> None: "reads out is the conditional distribution of the columns it does " "not condition" ) - requested = _resolve_requested_columns( - self.schema.columns, self.requested_columns - ) - if requested is None: - return - requested_names = {value for value in requested if isinstance(value, str)} - pinned = sorted(set(block.pinned_columns) & requested_names) - if pinned: - raise ValueError( - f"requested_columns lists pinned columns {pinned}; an equality " - "condition fixes the value, so reading it back returns only what " - "the request supplied" - ) def _validate_query_rows(self) -> None: """Check the observed rows a log-density request scores. The service scores the whole joint at each future position, so a row must carry every declared column and ``requested_columns`` must name - every readable one in schema order: a narrower projection would score a - different distribution than the caller believes it asked about. An - equality condition is the one exception, since it fixes its column's - value and the conditioning contract forbids reading it back. + every one of them in declared order: a narrower or reordered projection + would score a different distribution than the caller believes it asked + about. A ``condition`` block changes nothing here — its columns are + declared columns like any other — and omitting the field already + resolves to exactly this projection. Only two per-value failures are decidable without the deployment's own encoding and are checked here — a missing value and a non-finite number. @@ -643,16 +632,11 @@ def _validate_query_rows(self) -> None: ) if requested is None: return - pinned = set(() if self.condition is None else self.condition.pinned_columns) - readable = [name for name in declared if name not in pinned] - if requested != readable: - excused = ( - f", the pinned columns {sorted(pinned)} excepted" if pinned else "" - ) + if requested != declared: raise ValueError( "return_mode='log_prob' scores the whole joint, so " - "requested_columns must list every declared column in schema " - f"order{excused}: expected {readable}, got {requested}" + "requested_columns must list every declared column in declared " + f"order: expected {declared}, got {requested}" ) def to_payload(self) -> dict[str, Any]: diff --git a/tests/fixtures/condition_interval_response.json b/tests/fixtures/condition_interval_response.json index e9ab16e..f217332 100644 --- a/tests/fixtures/condition_interval_response.json +++ b/tests/fixtures/condition_interval_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/condition_log_prob_request.json b/tests/fixtures/condition_log_prob_request.json index 2407223..b6ef705 100644 --- a/tests/fixtures/condition_log_prob_request.json +++ b/tests/fixtures/condition_log_prob_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "condition", "return_mode": "log_prob", @@ -26,7 +26,7 @@ } ], "query_times": [2, 3], - "requested_columns": ["target"], + "requested_columns": ["driver", "target"], "condition": { "query_time_index": 1, "conditions": [ diff --git a/tests/fixtures/condition_log_prob_response.json b/tests/fixtures/condition_log_prob_response.json index 21d9f62..04513ea 100644 --- a/tests/fixtures/condition_log_prob_response.json +++ b/tests/fixtures/condition_log_prob_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", @@ -8,7 +8,7 @@ "return_mode": "log_prob", "outputs": { "query_times": [3], - "requested_columns": ["target"], + "requested_columns": ["driver", "target"], "mean": null, "samples": null, "quantiles": null, diff --git a/tests/fixtures/condition_mean_request.json b/tests/fixtures/condition_mean_request.json index ba7d9f4..f506b20 100644 --- a/tests/fixtures/condition_mean_request.json +++ b/tests/fixtures/condition_mean_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "condition", "return_mode": "mean", diff --git a/tests/fixtures/condition_mean_response.json b/tests/fixtures/condition_mean_response.json index 622745b..0c06ed6 100644 --- a/tests/fixtures/condition_mean_response.json +++ b/tests/fixtures/condition_mean_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/forecast_log_prob_request.json b/tests/fixtures/forecast_log_prob_request.json index f843ce4..6ba470f 100644 --- a/tests/fixtures/forecast_log_prob_request.json +++ b/tests/fixtures/forecast_log_prob_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "forecast", "return_mode": "log_prob", diff --git a/tests/fixtures/forecast_log_prob_response.json b/tests/fixtures/forecast_log_prob_response.json index f6ad94a..ab9ee00 100644 --- a/tests/fixtures/forecast_log_prob_response.json +++ b/tests/fixtures/forecast_log_prob_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/forecast_mean_request.json b/tests/fixtures/forecast_mean_request.json index c4b7ab7..17c29ff 100644 --- a/tests/fixtures/forecast_mean_request.json +++ b/tests/fixtures/forecast_mean_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "forecast", "return_mode": "mean", diff --git a/tests/fixtures/forecast_mean_response.json b/tests/fixtures/forecast_mean_response.json index d6aa094..d289237 100644 --- a/tests/fixtures/forecast_mean_response.json +++ b/tests/fixtures/forecast_mean_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/forecast_quantiles_request.json b/tests/fixtures/forecast_quantiles_request.json index 602ddc4..32442ba 100644 --- a/tests/fixtures/forecast_quantiles_request.json +++ b/tests/fixtures/forecast_quantiles_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "forecast", "return_mode": "quantiles", diff --git a/tests/fixtures/forecast_quantiles_response.json b/tests/fixtures/forecast_quantiles_response.json index aa8e9bd..cd49f12 100644 --- a/tests/fixtures/forecast_quantiles_response.json +++ b/tests/fixtures/forecast_quantiles_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/forecast_samples_request.json b/tests/fixtures/forecast_samples_request.json index 0f161c2..9203269 100644 --- a/tests/fixtures/forecast_samples_request.json +++ b/tests/fixtures/forecast_samples_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "forecast", "return_mode": "samples", diff --git a/tests/fixtures/forecast_samples_response.json b/tests/fixtures/forecast_samples_response.json index 60e60e1..29cff5c 100644 --- a/tests/fixtures/forecast_samples_response.json +++ b/tests/fixtures/forecast_samples_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/health_metadata.json b/tests/fixtures/health_metadata.json index 08d9aeb..804fdce 100644 --- a/tests/fixtures/health_metadata.json +++ b/tests/fixtures/health_metadata.json @@ -1,6 +1,6 @@ { "status": "ok", - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/input_size_exceeded_response.json b/tests/fixtures/input_size_exceeded_response.json index e244f5c..2e1c7c1 100644 --- a/tests/fixtures/input_size_exceeded_response.json +++ b/tests/fixtures/input_size_exceeded_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "errors": [ { "code": "INPUT_SIZE_EXCEEDED", diff --git a/tests/fixtures/model_version_mismatch_response.json b/tests/fixtures/model_version_mismatch_response.json index b5fa39e..07f5363 100644 --- a/tests/fixtures/model_version_mismatch_response.json +++ b/tests/fixtures/model_version_mismatch_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "errors": [ { "code": "MODEL_VERSION_MISMATCH", diff --git a/tests/fixtures/schema_version_mismatch_response.json b/tests/fixtures/schema_version_mismatch_response.json index ffb299b..d5cffa9 100644 --- a/tests/fixtures/schema_version_mismatch_response.json +++ b/tests/fixtures/schema_version_mismatch_response.json @@ -1,9 +1,9 @@ { - "schema_version": "v2", + "schema_version": "v3", "errors": [ { "code": "SCHEMA_VERSION_MISMATCH", - "message": "Unsupported schema_version: expected 'v3', got 'v2'", + "message": "Unsupported schema_version: expected 'v4', got 'v3'", "field": "schema_version", "detail": {} } diff --git a/tests/fixtures/validation_error_response.json b/tests/fixtures/validation_error_response.json index 7f12104..30d6160 100644 --- a/tests/fixtures/validation_error_response.json +++ b/tests/fixtures/validation_error_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "errors": [ { "code": "VALIDATION_ERROR", diff --git a/tests/test_cli.py b/tests/test_cli.py index 440447a..006392a 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -48,7 +48,7 @@ class FakeHealthClient: "deployment-id/predictionsUnstructured" ), deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -58,7 +58,7 @@ def health(self, *, cache: bool = False, refresh: bool = False) -> HealthMetadat del cache, refresh return HealthMetadata( status="ok", - schema_version="v3", + schema_version="v4", image_version="0.3.0", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", checkpoint_version="sdk-test", @@ -172,7 +172,7 @@ def test_predict_command_writes_response_file(monkeypatch, tmp_path: Path) -> No request_file.write_text( json.dumps( { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", } ), diff --git a/tests/test_condition_mode.py b/tests/test_condition_mode.py index 31d7deb..7a1e25c 100644 --- a/tests/test_condition_mode.py +++ b/tests/test_condition_mode.py @@ -97,7 +97,7 @@ def _health( """Build one advertisement without going through a transport.""" return HealthMetadata( status="ok", - schema_version="v3", + schema_version="v4", image_version="0.3.0", model_version=_MODEL_VERSION, checkpoint_version="sdk-test", @@ -195,14 +195,6 @@ def test_the_mode_and_the_block_must_agree() -> None: ("target",), "outside the 2 requested future positions", ), - ( - ConditionBlock( - query_time_index=0, - conditions=(EqualityCondition(column="driver", value=1.0),), - ), - ("driver", "target"), - "pinned columns", - ), ( ConditionBlock( query_time_index=0, @@ -249,6 +241,38 @@ def test_an_interval_column_may_still_be_read_out() -> None: assert request.to_payload()["requested_columns"] == ["driver", "target"] +def test_a_pinned_column_may_be_read_out_in_any_position() -> None: + """A projection is condition-agnostic, so a pin is nameable like any column. + + Its answer is the request's own value, which is what lets a scenario frame + be compared against an unconditioned one column for column. + """ + block = ConditionBlock( + query_time_index=0, + conditions=(EqualityCondition(column="driver", value=1.0),), + ) + + request = _request(block, requested_columns=("target", "driver")) + + assert request.to_payload()["requested_columns"] == ["target", "driver"] + + +def test_omitting_the_projection_states_nothing_on_the_wire() -> None: + """The default lives in the service, so the client must not invent one. + + A client-side default would be a second copy of the rule, free to drift from + the one the deployment actually applies. + """ + block = ConditionBlock( + query_time_index=0, + conditions=(EqualityCondition(column="driver", value=1.0),), + ) + + payload = _request(block, requested_columns=None).to_payload() + + assert "requested_columns" not in payload + + def test_the_payload_carries_the_block_the_service_parses() -> None: """The wire form names its position once and each condition names its kind.""" payload = build_forecast_payload( @@ -394,7 +418,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: assert isinstance(sample_count, int) start = sum(cast(int, earlier["n_samples"]) for earlier in self.payloads[:-1]) return { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": _MODEL_VERSION, "checkpoint_version": "sdk-test", @@ -432,7 +456,7 @@ def _health_payload( """Build one health advertisement as the service serializes it.""" return { "status": "ok", - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": _MODEL_VERSION, "checkpoint_version": "sdk-test", diff --git a/tests/test_configuration.py b/tests/test_configuration.py index 5ee8cab..3a6fe7b 100644 --- a/tests/test_configuration.py +++ b/tests/test_configuration.py @@ -47,7 +47,7 @@ class _HealthTransport: _METADATA: dict[str, object] = { "status": "ok", - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.yaml", "checkpoint_version": "yaml", @@ -147,7 +147,7 @@ def test_load_settings_layers_config_below_dotenv_and_environment( "datarobot_endpoint": "https://app.datarobot.com/api/v2", "datarobot_api_token": "yaml-token", "deployment_id": "yaml-deployment-id", - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.yaml", } }, @@ -191,7 +191,7 @@ def test_client_from_env_uses_transport_defaults_from_config( "datarobot_endpoint": "https://app.datarobot.com/api/v2", "datarobot_api_token": "yaml-token", "deployment_id": "yaml-deployment-id", - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.yaml", }, "transport": { diff --git a/tests/test_contract.py b/tests/test_contract.py index 90af77b..9d80ad3 100644 --- a/tests/test_contract.py +++ b/tests/test_contract.py @@ -76,7 +76,7 @@ def _health_metadata() -> dict[str, object]: """Health metadata.""" return { "status": "ok", - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -111,7 +111,7 @@ def test_package_identity_contract() -> None: assert DISTRIBUTION_NAME == "jointfm-client" assert IMPORT_NAMESPACE == "jointfm_client" assert FIRST_SUPPORTED_PYTHON_VERSION == "3.11" - assert SCHEMA_VERSION == "v3" + assert SCHEMA_VERSION == "v4" def test_package_version_matches_installed_distribution() -> None: @@ -267,7 +267,7 @@ def test_forecast_payload_matches_service_contract_without_mutating_inputs() -> ) assert payload == { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "quantiles", @@ -362,7 +362,7 @@ def test_dataframe_payload_matches_service_forecast_request_shape() -> None: ) assert payload == { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "mean", @@ -777,7 +777,7 @@ def test_health_and_response_models_parse_current_payloads() -> None: health = HealthMetadata.from_payload(_health_metadata()) response = ForecastResponse.from_payload( { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -840,7 +840,7 @@ def test_forecast_result_conversion_helpers_cover_mean_samples_and_quantiles() - """Forecast result conversion helpers cover mean samples and quantiles.""" mean_result = ForecastResponse.from_payload( { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -860,7 +860,7 @@ def test_forecast_result_conversion_helpers_cover_mean_samples_and_quantiles() - ) sample_result = ForecastResponse.from_payload( { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -883,7 +883,7 @@ def test_forecast_result_conversion_helpers_cover_mean_samples_and_quantiles() - ) quantile_result = ForecastResponse.from_payload( { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -957,7 +957,7 @@ def test_forecast_result_conversion_helpers_cover_mean_samples_and_quantiles() - def test_forecast_response_validates_request_scoped_shapes() -> None: """Forecast response validates request scoped shapes.""" request_payload = { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "samples", @@ -966,7 +966,7 @@ def test_forecast_response_validates_request_scoped_shapes() -> None: "n_samples": 3, } response_payload = { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -996,7 +996,7 @@ def test_forecast_response_raises_typed_error_for_success_payload_errors() -> No with pytest.raises(JointFMServiceError) as exc_info: ForecastResponse.from_payload( { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -1162,7 +1162,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: def _forecast_response_payload(*, return_mode: str) -> dict[str, object]: """Forecast response payload.""" return { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", diff --git a/tests/test_contract_models.py b/tests/test_contract_models.py index 9d5186b..af48ec8 100644 --- a/tests/test_contract_models.py +++ b/tests/test_contract_models.py @@ -76,7 +76,7 @@ def test_request_models_serialize_direct_payloads_without_mutating_inputs() -> N ) assert metadata.to_payload() == { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "quantiles", @@ -91,7 +91,7 @@ def test_request_models_serialize_direct_payloads_without_mutating_inputs() -> N "timezone": "UTC", } assert request.to_payload() == { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "quantiles", @@ -303,7 +303,7 @@ def test_response_models_reject_direct_validation_edges() -> None: def test_forecast_response_rejects_request_scoped_metadata_mismatches() -> None: """Forecast response rejects request scoped metadata mismatches.""" request_payload = { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "mean", @@ -354,7 +354,7 @@ def test_forecast_response_rejects_request_scoped_metadata_mismatches() -> None: def test_forecast_response_rejects_sample_bound_violations() -> None: """Forecast response rejects sample bound violations.""" request_payload = { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "samples", @@ -397,7 +397,7 @@ def test_forecast_response_rejects_sample_bound_violations() -> None: def _mean_response_payload() -> dict[str, Any]: """Mean response payload.""" return { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -419,7 +419,7 @@ def _mean_response_payload() -> dict[str, Any]: def _sample_response_payload() -> dict[str, Any]: """Sample response payload.""" return { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", diff --git a/tests/test_feature_importance.py b/tests/test_feature_importance.py index 3d6fe0a..c706175 100644 --- a/tests/test_feature_importance.py +++ b/tests/test_feature_importance.py @@ -47,7 +47,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: self.payloads.append(dict(payload)) samples = self.sample_batches[len(self.payloads) - 1] return { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": _MODEL_VERSION, "checkpoint_version": "sdk-test", diff --git a/tests/test_log_prob_mode.py b/tests/test_log_prob_mode.py index cdb6e01..9c0cd34 100644 --- a/tests/test_log_prob_mode.py +++ b/tests/test_log_prob_mode.py @@ -20,9 +20,10 @@ rows nothing else carries, and a projection narrower than the scored joint asks a different question than the caller believes they asked. The response carries a block with no column axis whose summary fields only repeat what ``values`` -already says. Under a condition the two meet: a pinned column leaves the -read-out set but still has to be supplied, because the service refuses a row -contradicting the pin rather than scoring it. +already says. Under a condition the two meet: a pinned column is projected and +supplied like any other, because the service refuses a row contradicting the +pin rather than scoring it, while the score itself stays the conditional +density of the columns the request does not condition. """ from __future__ import annotations @@ -170,18 +171,22 @@ def test_a_row_that_cannot_be_scored_is_refused_before_any_round_trip( def test_a_partial_projection_is_refused_because_it_scores_something_else() -> None: """Dropping a column would score a narrower joint than the caller believes.""" - with pytest.raises(ValueError, match="every declared column in schema order"): + with pytest.raises(ValueError, match="every declared column in declared order"): _request(requested_columns=("target",)) -def test_a_pinned_column_is_the_one_column_a_score_may_leave_out() -> None: - """Conditioning fixes the column, and the contract forbids reading it back.""" - pinned = _request( - requested_columns=("target",), +def test_a_score_may_leave_out_no_column_at_all_under_a_condition() -> None: + """A condition excuses nothing: the scorer covers the whole declared joint. + + Omitting the projection is the natural spelling, because the service's own + default already resolves to every declared column in declared order. + """ + stated = _request( + requested_columns=("driver", "target"), query_rows=_PINNED_QUERY_ROWS, condition=_PIN_DRIVER, ) - assert pinned.to_payload()["requested_columns"] == ["target"] + assert stated.to_payload()["requested_columns"] == ["driver", "target"] default_projection = _request( requested_columns=None, @@ -190,8 +195,12 @@ def test_a_pinned_column_is_the_one_column_a_score_may_leave_out() -> None: ) assert "requested_columns" not in default_projection.to_payload() - with pytest.raises(ValueError, match="lists pinned columns"): - _request(query_rows=_PINNED_QUERY_ROWS, condition=_PIN_DRIVER) + with pytest.raises(ValueError, match="every declared column in declared order"): + _request( + requested_columns=("target",), + query_rows=_PINNED_QUERY_ROWS, + condition=_PIN_DRIVER, + ) @pytest.mark.parametrize( @@ -201,7 +210,7 @@ def test_a_pinned_column_is_the_one_column_a_score_may_leave_out() -> None: ( "condition_log_prob_request", _PIN_DRIVER, - ("target",), + ("driver", "target"), _PINNED_QUERY_ROWS, ), ], @@ -329,7 +338,7 @@ def test_the_client_scores_under_a_condition( schema=_schema(), query_times=[2, 3], query_rows=list(_PINNED_QUERY_ROWS), - requested_columns=["target"], + requested_columns=["driver", "target"], model_version=_MODEL_VERSION, seed=7, condition=_PIN_DRIVER, diff --git a/tests/test_pool.py b/tests/test_pool.py index 837f19e..9de17bb 100644 --- a/tests/test_pool.py +++ b/tests/test_pool.py @@ -56,7 +56,7 @@ def _health( """Health.""" return { "status": "ok", - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": model_version, "checkpoint_version": checkpoint_version, @@ -134,7 +134,7 @@ def test_pool_retries_next_instance_on_470() -> None: pool = _pool(transport=_Transport(fail_ids=frozenset({"a"}))) assert pool.next_instance().deployment_id == "a" assert pool.next_instance().deployment_id == "b" - assert pool.post_json({"schema_version": "v3"}) == { + assert pool.post_json({"schema_version": "v4"}) == { "ok": True, "deployment_id": "b", } @@ -144,7 +144,7 @@ def test_pool_raises_when_all_instances_unavailable() -> None: """Pool raises when all instances unavailable.""" pool = _pool(transport=_Transport(fail_ids=frozenset({"a", "b"}))) with pytest.raises(JointFMHTTPStatusError, match="unavailable"): - pool.post_json({"schema_version": "v3"}) + pool.post_json({"schema_version": "v4"}) def test_pool_health_rejects_mismatch_and_aligns_sample_cap() -> None: @@ -275,7 +275,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: executor.submit( pool.post_json_to, pool.instance_at(index), - {"schema_version": "v3"}, + {"schema_version": "v4"}, ) for index in range(2) ] @@ -299,7 +299,7 @@ def test_pool_health_routes_only_reachable_peers() -> None: assert pool.instance_at(0).deployment_id == "b" assert pool.instance_at(1).deployment_id == "b" assert pool.next_instance().deployment_id == "b" - assert pool.post_json({"schema_version": "v3"})["deployment_id"] == "b" + assert pool.post_json({"schema_version": "v4"})["deployment_id"] == "b" def test_pool_failover_retries_health_excluded_peer() -> None: @@ -313,7 +313,7 @@ def test_pool_failover_retries_health_excluded_peer() -> None: assert pool.instance_at(0).deployment_id == "b" transport.fail_ids = frozenset({"b"}) - assert pool.post_json({"schema_version": "v3"})["deployment_id"] == "a" + assert pool.post_json({"schema_version": "v4"})["deployment_id"] == "a" assert pool.instance_at(0).deployment_id == "a" @@ -333,7 +333,7 @@ def test_pool_health_skips_incompatible_peer_when_another_matches_pin() -> None: assert metadata.model_version == pinned assert pool.instance_at(0).deployment_id == "a" assert pool.instance_at(1).deployment_id == "a" - assert pool.post_json({"schema_version": "v3"})["deployment_id"] == "a" + assert pool.post_json({"schema_version": "v4"})["deployment_id"] == "a" def test_pool_cooldown_restores_peer_after_transient_failure( @@ -345,7 +345,7 @@ def test_pool_cooldown_restores_peer_after_transient_failure( transport = _Transport(fail_ids=frozenset({"a"})) pool = _pool(transport=transport, peer_cooldown_seconds=10.0) - assert pool.post_json({"schema_version": "v3"})["deployment_id"] == "b" + assert pool.post_json({"schema_version": "v4"})["deployment_id"] == "b" assert pool.instance_at(0).deployment_id == "b" assert pool.instance_at(1).deployment_id == "b" @@ -355,4 +355,4 @@ def test_pool_cooldown_restores_peer_after_transient_failure( clock["now"] = 110.0 assert {pool.instance_at(i).deployment_id for i in range(2)} == {"a", "b"} - assert pool.post_json({"schema_version": "v3"})["deployment_id"] in {"a", "b"} + assert pool.post_json({"schema_version": "v4"})["deployment_id"] in {"a", "b"} diff --git a/tests/test_settings.py b/tests/test_settings.py index 22c74f3..13632c0 100644 --- a/tests/test_settings.py +++ b/tests/test_settings.py @@ -54,7 +54,7 @@ def _hosted_env(**overrides: str) -> dict[str, str]: DATAROBOT_ENDPOINT_ENV: "https://app.datarobot.com/api/v2/", DATAROBOT_API_TOKEN_ENV: "secret-token", JOINTFM_DEPLOYMENT_ID_ENV: "deployment-id", - JOINTFM_SCHEMA_VERSION_ENV: "v3", + JOINTFM_SCHEMA_VERSION_ENV: "v4", JOINTFM_MODEL_VERSION_ENV: "jointfm-inference:0.3.0+ckpt.sdk-test", } env.update(overrides) @@ -66,7 +66,7 @@ def test_load_settings_from_environment_with_deployment_id_builds_hosted_url() - settings = load_settings(env=_hosted_env(), dotenv_path=None) assert settings.datarobot_endpoint == "https://app.datarobot.com/api/v2" - assert settings.schema_version == "v3" + assert settings.schema_version == "v4" assert settings.model_version == "jointfm-inference:0.3.0+ckpt.sdk-test" assert settings.deployment_id == "deployment-id" assert settings.predict_url == ( @@ -82,7 +82,7 @@ def test_load_settings_with_local_service_base_url_builds_direct_urls() -> None: settings = load_settings( env={ JOINTFM_LOCAL_BASE_URL_ENV: "http://127.0.0.1:8080/", - JOINTFM_SCHEMA_VERSION_ENV: "v3", + JOINTFM_SCHEMA_VERSION_ENV: "v4", JOINTFM_MODEL_VERSION_ENV: "jointfm-inference:0.3.0+ckpt.local-test", }, dotenv_path=None, @@ -158,7 +158,7 @@ def test_load_settings_reads_dotenv_without_overriding_environment(tmp_path) -> "DATAROBOT_ENDPOINT=https://app.datarobot.com/api/v2", "DATAROBOT_API_TOKEN=file-token", "JOINTFM_DEPLOYMENT_ID=file-deployment-id", - "JOINTFM_SCHEMA_VERSION=v3", + "JOINTFM_SCHEMA_VERSION=v4", "JOINTFM_MODEL_VERSION=jointfm-inference:0.3.0+ckpt.sdk-test", ] ), @@ -209,7 +209,7 @@ def test_load_settings_rejects_missing_credentials_without_defaults() -> None: env={ DATAROBOT_API_TOKEN_ENV: "secret-token", JOINTFM_DEPLOYMENT_ID_ENV: "deployment-id", - JOINTFM_SCHEMA_VERSION_ENV: "v3", + JOINTFM_SCHEMA_VERSION_ENV: "v4", JOINTFM_MODEL_VERSION_ENV: "jointfm-inference:0.3.0+ckpt.sdk-test", }, dotenv_path=None, @@ -220,7 +220,7 @@ def test_load_settings_rejects_missing_credentials_without_defaults() -> None: env={ DATAROBOT_ENDPOINT_ENV: "https://app.datarobot.com/api/v2", JOINTFM_DEPLOYMENT_ID_ENV: "deployment-id", - JOINTFM_SCHEMA_VERSION_ENV: "v3", + JOINTFM_SCHEMA_VERSION_ENV: "v4", JOINTFM_MODEL_VERSION_ENV: "jointfm-inference:0.3.0+ckpt.sdk-test", }, dotenv_path=None, diff --git a/tests/test_surfaces.py b/tests/test_surfaces.py index 1d73b3d..bfd2a26 100644 --- a/tests/test_surfaces.py +++ b/tests/test_surfaces.py @@ -70,7 +70,7 @@ def test_hosted_surface_uses_datarobot_routes_and_auth_headers( health_url=predict_url, predict_url=predict_url, deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version=request_payload["model_version"], deployment_id="deployment-id", ) @@ -209,7 +209,7 @@ def test_hosted_surface_auto_discovers_model_version_when_settings_unpinned( health_url=predict_url, predict_url=predict_url, deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", deployment_id="deployment-id", ) assert settings.model_version is None diff --git a/tests/test_transport.py b/tests/test_transport.py index 441d7aa..c5fd1af 100644 --- a/tests/test_transport.py +++ b/tests/test_transport.py @@ -154,7 +154,7 @@ def _health_payload( """Health payload.""" return { "status": "ok", - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": model_version, "checkpoint_version": checkpoint_version, @@ -190,7 +190,7 @@ def _forecast_response_payload(*, return_mode: str = "mean") -> dict[str, object else None, } return { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", @@ -234,7 +234,7 @@ def test_transport_posts_json_with_headers_timeout_and_user_agent() -> None: session.mount("https://", adapter) result = transport.post_json( - "https://example.com/predict", {"schema_version": "v3"} + "https://example.com/predict", {"schema_version": "v4"} ) assert result == {"ok": True} @@ -248,7 +248,7 @@ def test_transport_posts_json_with_headers_timeout_and_user_agent() -> None: assert adapter.kwargs[0]["timeout"] == (1.5, 2.5) request_body = request.body assert isinstance(request_body, bytes) - assert json.loads(request_body.decode("utf-8")) == {"schema_version": "v3"} + assert json.loads(request_body.decode("utf-8")) == {"schema_version": "v4"} def test_transport_from_settings_attaches_hosted_auth_headers_and_closes_session() -> ( @@ -265,7 +265,7 @@ def test_transport_from_settings_attaches_hosted_auth_headers_and_closes_session "deployment-id/predictionsUnstructured" ), deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -293,7 +293,7 @@ def test_transport_from_local_settings_omits_hosted_auth_headers() -> None: health_url="http://127.0.0.1:8080/healthz", predict_url="http://127.0.0.1:8080/predict", deployment_selector="local_service", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.local-test", local_base_url="http://127.0.0.1:8080", ) @@ -319,7 +319,7 @@ def test_transport_retries_retryable_server_responses() -> None: retry_config=JointFMRetryConfig(max_attempts=2) ) - result = transport.post_json(_server_url(server), {"schema_version": "v3"}) + result = transport.post_json(_server_url(server), {"schema_version": "v4"}) assert result == {"ok": True} assert handler.request_count == 2 @@ -351,7 +351,7 @@ def test_transport_retries_html_bodied_gateway_errors() -> None: ), ) - result = transport.post_json(_server_url(server), {"schema_version": "v3"}) + result = transport.post_json(_server_url(server), {"schema_version": "v4"}) assert result == {"ok": True} assert handler.request_count == 2 @@ -380,7 +380,7 @@ def test_transport_raises_status_error_when_gateway_html_persists() -> None: ) with pytest.raises(JointFMHTTPStatusError) as exc_info: - transport.post_json(_server_url(server), {"schema_version": "v3"}) + transport.post_json(_server_url(server), {"schema_version": "v4"}) assert exc_info.value.status_code == HTTPStatus.BAD_GATEWAY assert "502 Bad Gateway" in exc_info.value.response_body_excerpt @@ -444,7 +444,7 @@ def request(self, *args: Any, **kwargs: Any) -> requests.Response: ) result = transport.post_json( - "https://example.com/predict", {"schema_version": "v3"} + "https://example.com/predict", {"schema_version": "v4"} ) assert result == {"ok": True} @@ -463,7 +463,7 @@ def test_status_error_carries_parsed_retry_after_seconds() -> None: ) with pytest.raises(JointFMHTTPStatusError) as exc_info: - transport.post_json(_server_url(server), {"schema_version": "v3"}) + transport.post_json(_server_url(server), {"schema_version": "v4"}) assert exc_info.value.retry_after_seconds == 0.5 assert handler.request_count == 1 @@ -481,7 +481,7 @@ def test_transport_does_not_retry_validation_errors() -> None: ) with pytest.raises(JointFMHTTPStatusError) as exc_info: - transport.post_json(_server_url(server), {"schema_version": "v3"}) + transport.post_json(_server_url(server), {"schema_version": "v4"}) assert exc_info.value.status_code == HTTPStatus.BAD_REQUEST assert exc_info.value.datarobot_request_id == "request-id-1" @@ -507,7 +507,7 @@ def test_transport_rejects_non_json_serializable_payloads() -> None: with pytest.raises(JointFMRequestEncodingError, match="JSON-serializable"): transport.post_json( "https://example.com/predict", - {"schema_version": "v3", "bad": object()}, + {"schema_version": "v4", "bad": object()}, ) @@ -572,14 +572,14 @@ def test_client_predict_uses_configured_transport_and_settings() -> None: health_url="https://app.datarobot.com/api/v2/deployments/deployment-id/healthz", predict_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) transport = RecordingTransport() client = JointFMClient(settings=settings, transport=transport) payload = { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", } @@ -598,7 +598,7 @@ def test_client_health_returns_typed_metadata_and_caches_only_when_requested() - health_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", predict_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -629,7 +629,7 @@ def test_client_health_instances_returns_one_entry_for_single_endpoint() -> None health_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", predict_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -691,7 +691,7 @@ def capture_transport( "DATAROBOT_ENDPOINT": "https://app.datarobot.com/api/v2", "DATAROBOT_API_TOKEN": "secret-token", "JOINTFM_DEPLOYMENT_ID": "deployment-id", - "JOINTFM_SCHEMA_VERSION": "v3", + "JOINTFM_SCHEMA_VERSION": "v4", "JOINTFM_MODEL_VERSION": "jointfm-inference:0.3.0+ckpt.sdk-test", }, dotenv_path=None, @@ -713,7 +713,7 @@ def test_client_health_rejects_cached_model_mismatch() -> None: health_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", predict_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -739,7 +739,7 @@ def test_client_hosted_health_posts_request_type_health_to_predict_url() -> None health_url=predict_url, predict_url=predict_url, deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -763,7 +763,7 @@ def test_client_local_health_keeps_get_request_to_healthz_route() -> None: health_url="http://127.0.0.1:8080/healthz", predict_url="http://127.0.0.1:8080/predict", deployment_selector="local_service", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", local_base_url="http://127.0.0.1:8080", ) @@ -805,7 +805,7 @@ def test_client_forecast_builds_payload_from_rows_and_returns_typed_response() - assert result.outputs.mean == ((12.0,),) assert transport.predict_url == "http://localhost:8080/predict" assert transport.payload == { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "forecast", "return_mode": "mean", @@ -980,7 +980,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: health_url="http://127.0.0.1:8080/healthz", predict_url="http://127.0.0.1:8080/predict", deployment_selector="local_service", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", local_base_url="http://127.0.0.1:8080", ) @@ -1032,13 +1032,13 @@ def test_client_predict_raises_typed_service_error_for_success_payload_errors() health_url="https://app.datarobot.com/api/v2/deployments/deployment-id/healthz", predict_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) transport = RecordingTransport() transport.predict_payload = { - "schema_version": "v3", + "schema_version": "v4", "errors": [ { "code": "VALIDATION_ERROR", @@ -1052,7 +1052,7 @@ def test_client_predict_raises_typed_service_error_for_success_payload_errors() with pytest.raises(JointFMServiceError) as exc_info: client.predict( { - "schema_version": "v3", + "schema_version": "v4", "model_version": settings.model_version, } ) @@ -1141,7 +1141,7 @@ def do_POST(self) -> None: payload = {"ok": True} else: payload = { - "schema_version": "v3", + "schema_version": "v4", "errors": [ { "code": "VALIDATION_ERROR", @@ -1225,7 +1225,7 @@ def _pool_settings(primary: str, backup: str) -> JointFMSettings: health_url=primary, predict_url=primary, deployment_selector="deployment_ids", - schema_version="v3", + schema_version="v4", instances=( JointFMInstanceSettings(deployment_id="primary-id", predict_url=primary), JointFMInstanceSettings(deployment_id="backup-id", predict_url=backup), @@ -1267,7 +1267,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: transport = PoolTransport() client = JointFMClient(settings=settings, transport=transport) payload = { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", } client.predict(payload) From 6d78b667fa282a82a2b6271d0f15af97de8f9c04 Mon Sep 17 00:00:00 2001 From: Stefan Hackmann Date: Wed, 23 Sep 2026 10:47:17 +0000 Subject: [PATCH 4/8] chore: Update the documented wire contract to schema version v4 --- .env.sample | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.env.sample b/.env.sample index e98b8c8..896f455 100644 --- a/.env.sample +++ b/.env.sample @@ -4,7 +4,7 @@ DATAROBOT_ENDPOINT=https://app.datarobot.com/api/v2 DATAROBOT_API_TOKEN= # JointFM compatibility metadata for the selected deployment. -JOINTFM_SCHEMA_VERSION=v3 +JOINTFM_SCHEMA_VERSION=v4 # Optional drift-detection pin; the SDK discovers the model version from /healthz when unset: # JOINTFM_MODEL_VERSION=jointfm-inference:0.2.0+ckpt.fin-2026-05-22 From 773bcbd888fae9c1dc7fe9309237bd38c83c9ca2 Mon Sep 17 00:00:00 2001 From: Stefan Hackmann Date: Wed, 23 Sep 2026 11:18:12 +0000 Subject: [PATCH 5/8] docs: Explaining log prob result --- notebooks/forecast_condition.ipynb | 2 +- src/jointfm_client/client.py | 3 ++- 2 files changed, 3 insertions(+), 2 deletions(-) diff --git a/notebooks/forecast_condition.ipynb b/notebooks/forecast_condition.ipynb index d3fc9f9..4fac89e 100644 --- a/notebooks/forecast_condition.ipynb +++ b/notebooks/forecast_condition.ipynb @@ -444,7 +444,7 @@ "## Scoring an outcome you already have\n", "Every read-out so far answers *what does the model expect*. `log_prob` asks the opposite — how plausible are values I already hold — and is the only return mode where the caller supplies the future instead of receiving it. Under a condition it scores those values against the conditional, which is how two candidate outcomes of the same what-if are compared.\n", "\n", - "Three rules follow from scoring a joint rather than a projection of one, and the client enforces all three before anything is sent. `query_rows` carries one observed row per entry of `query_times`, each with a value for every declared column. `requested_columns` must name every *readable* column in schema order — a narrower projection would score a different distribution than the caller means to ask about — and the pinned column is the one exception, since conditioning fixed its value. The deployment refuses a row that contradicts the condition, a pinned column given another value or a banded column outside its range, instead of scoring it.\n", + "Three rules follow from scoring a joint rather than a projection of one, and the client enforces all three before anything is sent. `query_rows` carries one observed row per entry of `query_times`, each with a value for every declared column. `requested_columns` must name every declared column in declared order — a narrower projection would score a different distribution than the caller means to ask about — and a condition excuses nothing, so omitting the field is the natural spelling. The deployment refuses a row that contradicts the condition, a pinned column given another value or a banded column outside its range, instead of scoring it.\n", "\n", "The scores below are log densities of whole rows, so the unit caveat from the plausibility section applies unchanged: read their difference, not their level." ] diff --git a/src/jointfm_client/client.py b/src/jointfm_client/client.py index c15b1f9..8a6a5bd 100644 --- a/src/jointfm_client/client.py +++ b/src/jointfm_client/client.py @@ -541,7 +541,8 @@ def forecast_log_prob( of ``query_times``, and the result is the model's log density of those values under its joint at that position. It therefore scores the whole joint, and ``requested_columns`` — when given at all — must name every - readable column in schema order. + declared column in declared order, which is what omitting it already + resolves to. With ``condition`` the score is taken under the conditional at the one future position the block names, and the service refuses rows that From f189d4052a4d8769a280f3875206b3a048d743e9 Mon Sep 17 00:00:00 2001 From: Stefan Hackmann Date: Wed, 23 Sep 2026 17:14:47 +0200 Subject: [PATCH 6/8] feat: Replace ConditionBlock with per-condition query_time_indices for schema v5 --- .env.sample | 2 +- README.md | 50 ++- config.sample.yaml | 2 +- docs/api-reference.md | 37 +- notebooks/forecast_condition.ipynb | 168 ++++++--- src/jointfm_client/__init__.py | 6 +- src/jointfm_client/adapters.py | 11 +- src/jointfm_client/client.py | 46 ++- src/jointfm_client/contract.py | 293 ++++++++------ .../fixtures/condition_interval_response.json | 7 +- .../fixtures/condition_log_prob_request.json | 20 +- .../fixtures/condition_log_prob_response.json | 18 +- tests/fixtures/condition_mean_request.json | 20 +- tests/fixtures/condition_mean_response.json | 7 +- tests/fixtures/forecast_log_prob_request.json | 2 +- .../fixtures/forecast_log_prob_response.json | 2 +- tests/fixtures/forecast_mean_request.json | 2 +- tests/fixtures/forecast_mean_response.json | 2 +- .../fixtures/forecast_quantiles_request.json | 2 +- .../fixtures/forecast_quantiles_response.json | 2 +- tests/fixtures/forecast_samples_request.json | 2 +- tests/fixtures/forecast_samples_response.json | 2 +- tests/fixtures/health_metadata.json | 2 +- .../input_size_exceeded_response.json | 2 +- .../model_version_mismatch_response.json | 2 +- .../schema_version_mismatch_response.json | 4 +- tests/fixtures/validation_error_response.json | 2 +- tests/test_cli.py | 6 +- tests/test_condition_mode.py | 356 +++++++++++------- tests/test_configuration.py | 6 +- tests/test_contract.py | 24 +- tests/test_contract_models.py | 12 +- tests/test_feature_importance.py | 2 +- tests/test_log_prob_mode.py | 17 +- tests/test_pool.py | 18 +- tests/test_settings.py | 12 +- tests/test_surfaces.py | 4 +- tests/test_transport.py | 58 +-- 38 files changed, 719 insertions(+), 511 deletions(-) diff --git a/.env.sample b/.env.sample index 896f455..167b6da 100644 --- a/.env.sample +++ b/.env.sample @@ -4,7 +4,7 @@ DATAROBOT_ENDPOINT=https://app.datarobot.com/api/v2 DATAROBOT_API_TOKEN= # JointFM compatibility metadata for the selected deployment. -JOINTFM_SCHEMA_VERSION=v4 +JOINTFM_SCHEMA_VERSION=v5 # Optional drift-detection pin; the SDK discovers the model version from /healthz when unset: # JOINTFM_MODEL_VERSION=jointfm-inference:0.2.0+ckpt.fin-2026-05-22 diff --git a/README.md b/README.md index 2b1e7f0..c795d61 100644 --- a/README.md +++ b/README.md @@ -10,7 +10,7 @@ The SDK targets the DataRobot-hosted unstructured prediction route and the same - Import namespace: `jointfm_client` - Supported Python: `>=3.11` - Current SDK package version: `0.8.0` -- Current JointFM service schema: `schema_version="v4"` +- Current JointFM service schema: `schema_version="v5"` The public API shape is a synchronous low-level `JointFMClient` with `health()`, `health_instances()`, and `predict(payload)` methods plus high-level `forecast(...)`, `forecast_mean(...)`, `forecast_samples(...)`, and `forecast_quantiles(...)` helpers. The SDK is not a proxy service; callers use it as a local Python library that talks to the hosted or local JointFM endpoint. @@ -50,7 +50,7 @@ Example deployment configuration: deployment: datarobot_endpoint: https://app.datarobot.com/api/v2 datarobot_api_token: - schema_version: v4 + schema_version: v5 deployment_id: # Optional model-version pin; the SDK discovers it from /healthz when unset: # model_version: jointfm-inference:0.3.0+ckpt.fin-2026-05-22 @@ -68,7 +68,7 @@ Equivalent `.env` deployment configuration: ```dotenv DATAROBOT_ENDPOINT=https://app.datarobot.com/api/v2 DATAROBOT_API_TOKEN= -JOINTFM_SCHEMA_VERSION=v4 +JOINTFM_SCHEMA_VERSION=v5 JOINTFM_DEPLOYMENT_ID= # Optional drift-detection pin; the SDK discovers the model version from /healthz when unset: # JOINTFM_MODEL_VERSION=jointfm-inference:0.3.0+ckpt.fin-2026-05-22 @@ -78,7 +78,7 @@ Equivalent local REST configuration for a service started from the `joint` repos ```dotenv JOINTFM_LOCAL_BASE_URL=http://127.0.0.1:8080 -JOINTFM_SCHEMA_VERSION=v4 +JOINTFM_SCHEMA_VERSION=v5 # Optional drift-detection pin; the SDK discovers the model version from /healthz when unset: # JOINTFM_MODEL_VERSION=jointfm-inference:0.3.0+ckpt.fin_i504_o63_f0_t10_h16l16_mam7_af_t3r1_cnn_k3l4_hpst_h16l2_studentt_m4cr2df8skew ``` @@ -167,51 +167,47 @@ The bootstrap helper resolves the nearest src-layout Python project root, switch The current forecast request contract is: -- `schema_version`: exactly `"v4"`, configured as `JOINTFM_SCHEMA_VERSION` for `from_env()` clients +- `schema_version`: exactly `"v5"`, configured as `JOINTFM_SCHEMA_VERSION` for `from_env()` clients - `model_version`: exact model version advertised by `/healthz` or otherwise selected by the caller. Optional for `from_env()` clients: when `JOINTFM_MODEL_VERSION` is unset the SDK reads it from `/healthz` on first use; when set it acts as a drift-detection pin -- `query_mode`: `"forecast"` for the unconditional forecast, or `"condition"` for a conditional query at one future position; the high-level helpers set it from whether a `condition` block was passed +- `query_mode`: `"forecast"` for the unconditional forecast, or `"condition"` for the forecast given conditions on some columns; the high-level helpers set it from whether a `condition` was passed - `return_mode`: one of `"mean"`, `"samples"`, or `"quantiles"` - `time_index_mode`: one of `"ordinal"`, `"continuous_float"`, or `"absolute_datetime"` - `time_column`: required for `"absolute_datetime"`, and used for ordered ordinal or continuous histories when supplied - `query_times`: non-empty future forecast times only - `requested_columns`: optional column names or integer column indices, with duplicates rejected - `n_samples`: positive sample count for sampled forecasts and quantile estimation. When `return_mode="samples"` exceeds the `max_sample_count` advertised by the deployment's health metadata, `forecast_samples(...)` splits the request into capped prediction batches up front and returns one merged `SampleForecastResult`. -- `condition`: required with `query_mode="condition"` and forbidden otherwise. A `ConditionBlock` naming one future position by its index into `query_times` and one condition per column, see below. +- `condition`: required with `query_mode="condition"` and forbidden otherwise. One `EqualityCondition` or `IntervalCondition`, or a list of them, each covering the future positions its `query_time_indices` names; sent on the wire as the `conditions` list, see below. ### Conditional Queries -The `condition` query mode asks for the model's joint distribution at one future position *given* something about some of its columns at that same position. Conditioning relates columns to each other within one position and never across horizons, so the block names the position once and the response describes that position alone: `outputs.query_times` carries exactly one entry however many `query_times` the request listed. +The `condition` query mode asks for the forecast *given* something about some of its columns. Each condition names the future positions it covers by index into `query_times` through `query_time_indices`, and covers every position when that is `None`. Two conditions on the same column must cover disjoint positions; conditions on different columns may share positions, which is how one position mixes both kinds: -A column carries at most one condition, of either kind: +- `EqualityCondition(column, value, query_time_indices=None)` pins the column to a finite value. It stays nameable in `requested_columns` and reads back the value the request supplied, so a scenario answer lines up column for column with an unconditioned one. +- `IntervalCondition(column, lower=None, upper=None, query_time_indices=None)` confines the column to a range; `None` leaves that side open, and at least one side must be bounded. The column stays readable, and what comes back is its distribution inside the range. -- `EqualityCondition(column, value)` pins the column to a finite value. It stays nameable in `requested_columns` and reads back the value the request supplied, so a scenario answer lines up column for column with an unconditioned one. -- `IntervalCondition(column, lower=None, upper=None)` confines the column to a range; `None` leaves that side open, and at least one side must be bounded. The column stays readable, and what comes back is its distribution inside the range. - -Every column without a condition is a read-out column, and at least one must remain. `requested_columns` chooses what the response carries independently of that, defaulting to every declared column in declared order. Pass the block to `forecast(...)`, `forecast_mean(...)`, `forecast_samples(...)`, or `forecast_quantiles(...)`: +The response answers every entry of `query_times`. A position no condition covers carries the unconditioned forecast, and that is exact: the model draws each future position from its own joint, independently of the others, so a condition relates columns to each other within a position and never reaches another one. At every covered position at least one column must stay unconditioned, because what the position reads out is the conditional distribution of those columns. `requested_columns` chooses what the response carries independently of that, defaulting to every declared column in declared order. Pass one condition or a list of them to `forecast(...)`, `forecast_mean(...)`, `forecast_samples(...)`, `forecast_quantiles(...)`, or `forecast_log_prob(...)`: ```python -from jointfm_client import ConditionBlock, EqualityCondition, IntervalCondition +from jointfm_client import EqualityCondition, IntervalCondition -block = ConditionBlock( - query_time_index=0, - conditions=[ - EqualityCondition(column="equity_index_level", value=4780.0), - IntervalCondition(column="treasury_10y_yield", lower=0.041, upper=0.045), - ], -) result = client.forecast_mean( history, query_times=query_times, requested_columns=["portfolio_nav", "realized_volatility"], columns=plan.columns, - condition=block, + condition=[ + EqualityCondition(column="equity_index_level", value=4780.0), + IntervalCondition( + column="treasury_10y_yield", lower=0.041, upper=0.045, query_time_indices=[0] + ), + ], ) print(result.plausibility) ``` -Whether a deployment can condition depends on the mounted checkpoint's head. `/healthz` advertises `condition` in `supported_query_modes` and the kinds it answers in `supported_condition_kinds` (empty when the mode is absent). The client checks that advertisement before sending, so a deployment that cannot condition is refused with `UnsupportedServiceContractError` rather than after a paid round trip; `require_condition_support(metadata, block)` exposes the same check. +Whether a deployment can condition depends on the mounted checkpoint's head. `/healthz` advertises `condition` in `supported_query_modes` and the kinds it answers in `supported_condition_kinds` (empty when the mode is absent). The client checks that advertisement before sending, so a deployment that cannot condition is refused with `UnsupportedServiceContractError` rather than after a paid round trip; `require_condition_support(metadata, condition)` exposes the same check. -A condition response carries a `plausibility` block: `equality_log_density` is the log density the model assigns to the pinned values and `region_log_probability` the log probability it gives the interval region, each `None` when the request carried no condition of that kind. They separate a confident answer from one conditioned on something the model finds implausible; the service reports them and never refuses on them. `diagnostics.condition_draws` counts the draws behind sampled outputs, and `diagnostics.interval_estimator` reports the numerical accounting (`points`, `effective_sample_size`) when more than one column carries an interval and the region probability had to be estimated. +A condition response carries a `plausibility` block: `equality_log_density` is the log density the model assigns to the pinned values and `region_log_probability` the log probability it gives the interval region, each `None` when the request carried no condition of that kind. Over several covered positions each is the sum of the per-position values, exact because positions are independent. They separate a confident answer from one conditioned on something the model finds implausible; the service reports them and never refuses on them. `diagnostics.condition_draws` counts the draws behind sampled outputs, and `diagnostics.interval_estimator` reports the numerical accounting (`points`, `effective_sample_size`) when some position bounds more than one column and its region probability had to be estimated; across several estimated positions it describes the one with the smallest effective sample size. Column descriptors support the server fields `name`, `modality`, `role`, `nullable`, `vocabulary_size`, `level_count`, `mapping`, `lower_bound`, `upper_bound`, `time_value_kind`, `time_value_scale_seconds`, `time_value_use_local_normalized_time`, `time_value_calendar_id`, and `time_value_timezone`. @@ -227,7 +223,7 @@ Successful forecast responses preserve `schema_version`, `image_version`, `model ```json { - "schema_version": "v4", + "schema_version": "v5", "errors": [ { "code": "VALIDATION_ERROR", @@ -242,7 +238,7 @@ Known error codes are `VALIDATION_ERROR`, `UNSUPPORTED_HEAD_QUERY_COMBINATION`, ## Compatibility Policy -The SDK supports only `schema_version="v4"`. `validate_service_metadata()` checks `/healthz` metadata and raises typed compatibility errors before prediction if the service advertises a different schema, an unexpected model version, mode capabilities outside the recorded service contract, or an unsupported `decoding_strategy`. Return modes and time-index modes must match the SDK's lists exactly. Query modes and condition kinds are derived by the service from the mounted head, so a deployment may advertise fewer of them than the SDK knows; it must advertise at least one query mode and nothing the SDK does not know. +The SDK supports only `schema_version="v5"`. `validate_service_metadata()` checks `/healthz` metadata and raises typed compatibility errors before prediction if the service advertises a different schema, an unexpected model version, mode capabilities outside the recorded service contract, or an unsupported `decoding_strategy`. Return modes and time-index modes must match the SDK's lists exactly. Query modes and condition kinds are derived by the service from the mounted head, so a deployment may advertise fewer of them than the SDK knows; it must advertise at least one query mode and nothing the SDK does not know. Callers should pass an expected `model_version` when they already know which deployment artifact they intend to use. A mismatch is treated as a hard compatibility error rather than silently downgrading, guessing, or retrying another model. @@ -284,7 +280,7 @@ Create `.env` from `.env.sample` or set the same values in your shell. A hosted ```dotenv DATAROBOT_ENDPOINT=https://app.datarobot.com/api/v2 DATAROBOT_API_TOKEN= -JOINTFM_SCHEMA_VERSION=v4 +JOINTFM_SCHEMA_VERSION=v5 JOINTFM_DEPLOYMENT_ID= # Or: JOINTFM_DEPLOYMENT_IDS=chevron-id,research-id # Optional drift-detection pin; the SDK discovers the model version from /healthz when unset: diff --git a/config.sample.yaml b/config.sample.yaml index a49d28e..dd45a56 100644 --- a/config.sample.yaml +++ b/config.sample.yaml @@ -52,7 +52,7 @@ transport: - X-DataRobot-Execution-ID user_agent_header: User-Agent forecast: - schema_version: v4 + schema_version: v5 query_mode: forecast return_mode: mean time_index_mode: ordinal diff --git a/docs/api-reference.md b/docs/api-reference.md index efc7f82..ced93bf 100644 --- a/docs/api-reference.md +++ b/docs/api-reference.md @@ -6,7 +6,7 @@ This reference covers the supported public Python surface exported by `jointfm_c | Name | Purpose | | --- | --- | -| `JointFMClient` | Synchronous client for hosted or local JointFM endpoints. Use `from_env()` for `.env` and `config.yaml` backed hosted settings, `health()` for consensus typed service metadata, `health_instances()` for per-deployment probe results and pooled sample topology, `predict(payload)` for low-level JSON prediction, `forecast(...)` for validated tabular forecasts, and the `forecast_mean(...)`, `forecast_samples(...)`, `forecast_quantiles(...)`, and `forecast_log_prob(...)` convenience methods for typed forecast results. `forecast_log_prob(...)` is the one that asks about values the caller already holds: it takes `query_rows`, one observed row per entry of `query_times` carrying every declared column, and returns their log density under the model's joint. Each forecast method accepts `condition=ConditionBlock(...)`, which switches the request to the `condition` query mode and, before anything is sent, checks that the deployment's health metadata advertises the mode and every condition kind the block uses. `health()` probes `GET /healthz` for local deployments and POSTs `{"request_type": "health"}` to `predict_url` for hosted DataRobot deployments because the DataRobot deployment gateway only proxies the unstructured prediction route. `feature_importance(...)` runs permutation feature importance: one baseline `forecast_samples` call plus one per shuffled feature column, returning a list of `{"feature", "mean", "distance"}` dicts, each holding that feature's absolute forecast-mean shift and centered squared 2-Wasserstein distance indexed by target and horizon. | +| `JointFMClient` | Synchronous client for hosted or local JointFM endpoints. Use `from_env()` for `.env` and `config.yaml` backed hosted settings, `health()` for consensus typed service metadata, `health_instances()` for per-deployment probe results and pooled sample topology, `predict(payload)` for low-level JSON prediction, `forecast(...)` for validated tabular forecasts, and the `forecast_mean(...)`, `forecast_samples(...)`, `forecast_quantiles(...)`, and `forecast_log_prob(...)` convenience methods for typed forecast results. `forecast_log_prob(...)` is the one that asks about values the caller already holds: it takes `query_rows`, one observed row per entry of `query_times` carrying every declared column, and returns their log density under the model's joint. Each forecast method accepts `condition=`, one `EqualityCondition` or `IntervalCondition` or a list of them, which switches the request to the `condition` query mode and, before anything is sent, checks that the deployment's health metadata advertises the mode and every condition kind the conditions use. `health()` probes `GET /healthz` for local deployments and POSTs `{"request_type": "health"}` to `predict_url` for hosted DataRobot deployments because the DataRobot deployment gateway only proxies the unstructured prediction route. `feature_importance(...)` runs permutation feature importance: one baseline `forecast_samples` call plus one per shuffled feature column, returning a list of `{"feature", "mean", "distance"}` dicts, each holding that feature's absolute forecast-mean shift and centered squared 2-Wasserstein distance indexed by target and horizon. | `JointFMClient.from_env()` loads `config.yaml`, optional `.env` values, and process environment variables. `JointFMClient.health(cache=True)` caches health metadata only when requested. `JointFMClient.health_instances()` returns the same probe as a `HealthInstances` object: one `InstanceHealth` per configured deployment (including failures), `max_sample_count` as the sum of reachable caps (overall parallel capacity), and `topology` / `topology_label` grouping those caps (unavailable peers are listed but excluded from the sum and topology). Each endpoint's health payload describes only that endpoint; the client aggregates by calling each configured peer. `health()` still exposes the minimum reachable `max_sample_count`, which is the sample-batch cap used by forecast helpers. `JointFMClient.predict(payload)` requires `payload["model_version"]`; high-level forecast helpers resolve the configured model version when the caller does not pass one explicitly. When `forecast_samples(...)` requests an explicit `n_samples`, the client learns the deployment's `max_sample_count` from health metadata before the first prediction, splits oversized requests into capped prediction batches, and returns one merged `SampleForecastResult`. Clients configured without a reachable health route fall back to discovering the cap from the structured service error. @@ -17,12 +17,12 @@ This reference covers the supported public Python surface exported by `jointfm_c | `ColumnSpec` | Describes one modeled request column. Fields are `name`, `modality`, `role`, `nullable`, `vocabulary_size`, `level_count`, `mapping`, `lower_bound`, `upper_bound`, `time_value_kind`, `time_value_scale_seconds`, `time_value_use_local_normalized_time`, `time_value_calendar_id`, and `time_value_timezone`. | | `DataFrameSchema` | Describes tabular history layout. Fields are `columns`, `time_index_mode`, `time_column`, `time_scale_seconds`, `use_local_normalized_time`, `calendar_id`, and `timezone`. | | `ForecastRequestMetadata` | Holds `schema_version`, `model_version`, `query_mode` (`forecast` or `condition`), and `return_mode` for one forecast request. | -| `ForecastRequest` | Validated request object that combines metadata, schema, history rows, query times, requested columns, sample or quantile controls, `seed`, an optional `condition` block, and the optional `query_rows` a scored request supplies, then emits a JSON-compatible payload with `to_payload()`. `query_mode="condition"` and a `condition` block must appear together, as must `return_mode="log_prob"` and `query_rows`, and both are validated against the request's own schema and `query_times` before any round trip. | -| `EqualityCondition` | One column of the conditioned position pinned to a finite `value`. The column may be named in `requested_columns` like any other, and reads back the pinned value the request supplied. | -| `IntervalCondition` | One column of the conditioned position confined to `[lower, upper]`; `None` leaves a side open, at least one side must be bounded, and `lower < upper`. The column stays readable and the response describes its distribution inside the range. | -| `ConditionBlock` | Every condition of one request: `query_time_index` into the request's `query_times` and a sequence of `EqualityCondition` or `IntervalCondition` values, at most one per column. `kinds`, `pinned_columns`, and `conditioned_columns` expose what the capability gate and the request validation need. | +| `ForecastRequest` | Validated request object that combines metadata, schema, history rows, query times, requested columns, sample or quantile controls, `seed`, an optional `condition` (one condition or a list), and the optional `query_rows` a scored request supplies, then emits a JSON-compatible payload with `to_payload()`. `query_mode="condition"` and a `condition` must appear together, as must `return_mode="log_prob"` and `query_rows`, and both are validated against the request's own schema and `query_times` before any round trip. | +| `EqualityCondition` | One column pinned to a finite `value` at the positions `query_time_indices` names (indices into `query_times`, `None` for every position). The column may be named in `requested_columns` like any other, and reads back the pinned value the request supplied. | +| `IntervalCondition` | One column confined to `[lower, upper]` at the positions `query_time_indices` names (`None` for every position); `None` leaves a side open, at least one side must be bounded, and `lower < upper`. The column stays readable and the response describes its distribution inside the range. | +| `Condition` | Type alias for `EqualityCondition \| IntervalCondition`. Every `condition=` parameter takes one of them or a sequence of them. | | `ConditionPlausibility` | What the model thinks of the conditions it was given: `equality_log_density` of the pinned values and `region_log_probability` of the interval region, each `None` when the request carried no condition of that kind. Reported by the service and never refused on. | -| `IntervalEstimator` | Numerical accounting (`points`, `effective_sample_size`) behind a region probability that had to be estimated, which happens when more than one column carries an interval condition. | +| `IntervalEstimator` | Numerical accounting (`points`, `effective_sample_size`) behind a region probability that had to be estimated, which happens when some position bounds more than one column. Across several estimated positions it describes the one with the smallest effective sample size. | | `HealthMetadata` | Typed service-health payload with service status, schema and model versions, checkpoint metadata, device, head, `decoding_strategy`, advertised query modes, `supported_condition_kinds` (empty when the deployment cannot condition), return modes, time-index modes, time-index encoding, `max_sample_count`, and an optional `data_generation` block carrying advertised capacity limits. The container exposes it on `GET /healthz` for direct local access and as the response to `POST {"request_type": "health"}` on the unstructured prediction route for DataRobot-hosted deployments. Each endpoint reports only its own capabilities. | | `InstanceHealth` | One configured deployment's probe outcome: `deployment_id`, optional `metadata` (`HealthMetadata` when reachable), and optional `error` when the peer was skipped. | | `HealthInstances` | Client aggregation of `health_instances()`: `instances` (one `InstanceHealth` per configured ID), `max_sample_count` (sum of reachable caps = overall parallel capacity), `topology` as `(count, cap)` pairs sorted by descending cap, and `topology_label` such as `2x5000` or `1x7000, 1x3000`. Unavailable peers stay in `instances` but are omitted from the sum and topology. | @@ -80,9 +80,9 @@ All SDK-specific exceptions inherit from `JointFMError`. | `JointFMHTTPStatusError` | The service returns an HTTP error status. | | `JointFMServiceError` | A response body contains non-empty JointFM `errors`, including the case where HTTP status unexpectedly succeeded. | | `JointFMCompatibilityError` | Base class for fail-fast service compatibility failures. | -| `UnsupportedSchemaVersionError` | The service or response advertises a schema version other than `v4`. | +| `UnsupportedSchemaVersionError` | The service or response advertises a schema version other than `v5`. | | `UnsupportedModelVersionError` | The service or response model version differs from the configured or requested version. | -| `UnsupportedServiceContractError` | The service-health payload advertises mode capabilities or a `decoding_strategy` outside the recorded service contract, or a condition request targets a deployment that does not advertise the `condition` mode or one of the block's condition kinds. | +| `UnsupportedServiceContractError` | The service-health payload advertises mode capabilities or a `decoding_strategy` outside the recorded service contract, or a condition request targets a deployment that does not advertise the `condition` mode or one of the request's condition kinds. | ## Public Functions @@ -106,7 +106,8 @@ All SDK-specific exceptions inherit from `JointFMError`. | `build_datarobot_prediction_headers(api_token)` | Build hosted prediction headers: bearer authorization, broad accept header, and JSON content type. | | `build_forecast_payload(...)` | Build a validated JSON-compatible forecast payload from explicit schema, history rows, query times, and return-mode controls. | | `validate_service_metadata(metadata, expected_model_version=None)` | Validate the service-health metadata against the supported schema version, the expected model when supplied, advertised mode capabilities, and a supported `decoding_strategy`. Return and time-index modes must match the SDK's lists exactly; `supported_query_modes` (non-empty) and `supported_condition_kinds` (may be empty) must be subsets of what the SDK knows, because the service derives them from the mounted head. | -| `require_condition_support(metadata, block)` | Raise `UnsupportedServiceContractError` when `HealthMetadata` does not advertise the `condition` query mode or one of the kinds the `ConditionBlock` uses. The forecast helpers call it before sending a condition request. | +| `require_condition_support(metadata, condition)` | Raise `UnsupportedServiceContractError` when `HealthMetadata` does not advertise the `condition` query mode or one of the kinds `condition` uses. The forecast helpers call it before sending a condition request. | +| `resolve_conditions(condition)` | Normalize one condition or a sequence of them into a tuple, raising `ValueError` for an empty sequence, an entry that is not a condition, or two conditions on one column at a shared position (`query_time_indices=None` overlaps every other condition on its column). | | `infer_column_specs_from_dataframe(frame, ...)` | Infer ordered `ColumnSpec` objects from a pandas `DataFrame` and explicit role, modality, mapping, nullability, time-value, and bounds hints. | | `dataframe_to_history_rows(frame, schema)` | Convert a pandas `DataFrame` into server-compatible `history_rows`. | | `arrays_to_history_rows(values, columns=..., ...)` | Convert a two-dimensional NumPy-like array plus column metadata into `history_rows`. | @@ -126,7 +127,7 @@ All SDK-specific exceptions inherit from `JointFMError`. | --- | --- | --- | | `DATAROBOT_ENDPOINT` | Hosted calls | HTTPS DataRobot API v2 endpoint, normalized without a trailing slash and required to end in `/api/v2`. | | `DATAROBOT_API_TOKEN` | Hosted calls | Non-empty, whitespace-free API token used in the hosted bearer authorization header. | -| `JOINTFM_SCHEMA_VERSION` | Hosted calls | Request schema pin. The SDK supports only `v4`. | +| `JOINTFM_SCHEMA_VERSION` | Hosted calls | Request schema pin. The SDK supports only `v5`. | | `JOINTFM_MODEL_VERSION` | Hosted calls | Exact JointFM deployment model version expected from the service-health payload and prediction responses. | | `JOINTFM_DEPLOYMENT_ID` | One selector | Deployment ID used to build hosted health and prediction URLs. | | `JOINTFM_DEPLOYMENT_IDS` | One selector | Comma-separated hosted deployment IDs for round-robin load balancing (at least two unique IDs). Mutually exclusive with other selectors. Peers must share `model_version` and `checkpoint_version`. `health()` uses the minimum reachable `max_sample_count` as the sample-batch cap; `health_instances()` sums reachable caps for overall parallel capacity and reports topology. | @@ -157,9 +158,9 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | Field | Required | Description | | --- | --- | --- | | `request_type` | Optional | One of `"predict"` (default) or `"health"`. Forecast requests omit this field or set it to `"predict"`. | -| `schema_version` | Yes | Must be `"v4"`. | +| `schema_version` | Yes | Must be `"v5"`. | | `model_version` | Yes | Exact deployed model version expected by the caller. | -| `query_mode` | Yes | `"forecast"` for the unconditional forecast or `"condition"` for a conditional query at one future position. | +| `query_mode` | Yes | `"forecast"` for the unconditional forecast or `"condition"` for the forecast given the request's `conditions`. | | `return_mode` | Yes | One of `"mean"`, `"samples"`, `"quantiles"`, or `"log_prob"`, covered by the `forecast_mean`, `forecast_samples`, `forecast_quantiles`, and `forecast_log_prob` helpers. | | `time_index_mode` | Yes | One of `"ordinal"`, `"continuous_float"`, or `"absolute_datetime"`. | | `columns` | Yes | Non-empty array of column descriptors for modeled columns. | @@ -171,7 +172,7 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | `n_samples` | Samples and quantiles controls | Positive sample count when sampling controls are needed. Oversized sample forecasts are batched automatically against the cap advertised in health metadata. | | `quantiles` | Quantiles mode | Quantile levels in `(0, 1)`, required for `return_mode="quantiles"`. | | `seed` | Optional | Integer random seed for reproducible stochastic outputs. | -| `condition` | With `query_mode="condition"` | Object with `query_time_index` (index into `query_times`) and `conditions`, a list of `{"column", "kind": "equality", "value"}` or `{"column", "kind": "interval", "lower", "upper"}` entries with `null` for an open bound. At most one condition per column and at least one column left unconditioned. Forbidden with any other query mode. | +| `conditions` | With `query_mode="condition"` | Non-empty list of `{"column", "kind": "equality", "value", "query_time_indices"}` or `{"column", "kind": "interval", "lower", "upper", "query_time_indices"}` entries, `null` marking an open bound. `query_time_indices` is a list of distinct indices into `query_times`, or `null` for every position. Two conditions on one column must cover disjoint positions, and every covered position must leave at least one column unconditioned. Forbidden with any other query mode. | | `time_scale_seconds` | Optional | Positive scale for continuous time indexes. | | `use_local_normalized_time` | Optional | Whether the service should use local normalized time features. | | `calendar_id` | Optional | Calendar identifier, defaulting to `pandas-default`. | @@ -200,14 +201,14 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | Field | Description | | --- | --- | -| `schema_version` | Response schema, expected to be `"v4"`. | +| `schema_version` | Response schema, expected to be `"v5"`. | | `image_version` | Service image version that produced the response. | | `model_version` | Model version that produced the response. | | `checkpoint_version` | Checkpoint version that produced the response. | | `head` | Forecast head used by the service. | | `query_mode` | Response query mode, matching the request: `"forecast"` or `"condition"`. | | `return_mode` | Response return mode matching the request. | -| `outputs.query_times` | Forecast horizon values preserved from the request. A condition response carries only the conditioned position. | +| `outputs.query_times` | Forecast horizon values preserved from the request, on condition responses too: a position no condition covers carries the unconditioned forecast. | | `outputs.requested_columns` | Output columns in response order. | | `outputs.mean` | Mean values with axis order `(horizon, column)` when `return_mode="mean"`. | | `outputs.samples` | Sample values with axis order `(sample, horizon, column)` when `return_mode="samples"`. | @@ -218,7 +219,7 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | `outputs.log_prob` | `log_prob` responses only: per-horizon `values` and `nll_values` with their `total`, `mean`, `nll_total`, and `nll_mean` summaries. | | `diagnostics.condition_draws` | Condition responses only: number of draws behind sampled outputs. Batched sample requests report the merged count. | | `diagnostics.interval_estimator` | Condition responses only, and only when more than one column carries an interval: `points` and `effective_sample_size` of the numerical region-probability estimate. | -| `plausibility` | `null` on forecast responses. On condition responses an object with `equality_log_density` (log density of the pinned values) and `region_log_probability` (log probability of the interval region), each `null` when the request carried no condition of that kind. | +| `plausibility` | `null` on forecast responses. On condition responses an object with `equality_log_density` (log density of the pinned values) and `region_log_probability` (log probability of the interval region), each `null` when the request carried no condition of that kind. Over several covered positions each is the sum of the per-position values. | | `errors` | Structured service errors. Non-empty arrays raise typed SDK exceptions. | ### Health Metadata @@ -226,7 +227,7 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | Field | Description | | --- | --- | | `status` | Service status string. | -| `schema_version` | Advertised schema version. The SDK requires `v4`. | +| `schema_version` | Advertised schema version. The SDK requires `v5`. | | `image_version` | Running service image version. | | `model_version` | Running model version. | | `checkpoint_version` | Loaded checkpoint version. | @@ -235,7 +236,7 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | `head` | Active forecast head. | | `decoding_strategy` | Horizon decoding mode advertised by the mounted model. Must be one of `SUPPORTED_DECODING_STRATEGIES`: `parallel_dense`, `parallel_scalable`, or `autoregressive`. Parallel strategies decode every horizon in one pass; `autoregressive` rolls horizons sequentially. | | `supported_query_modes` | Non-empty subset of the SDK's query modes (`forecast`, `condition`); the service derives it from the mounted head. | -| `supported_condition_kinds` | Subset of the SDK's condition kinds (`equality`, `interval`), empty when `condition` is not advertised. Condition requests are refused locally when the block uses a kind that is missing here. | +| `supported_condition_kinds` | Subset of the SDK's condition kinds (`equality`, `interval`), empty when `condition` is not advertised. Condition requests are refused locally when a condition uses a kind that is missing here. | | `supported_return_modes` | Must match the SDK's return modes (`mean`, `samples`, `quantiles`, `log_prob`). | | `supported_time_index_modes` | Must match the SDK's time-index modes. | | `time_index_encoding` | Time-index encoding advertised by the service. | diff --git a/notebooks/forecast_condition.ipynb b/notebooks/forecast_condition.ipynb index 4fac89e..7a13611 100644 --- a/notebooks/forecast_condition.ipynb +++ b/notebooks/forecast_condition.ipynb @@ -49,14 +49,14 @@ "# Conditional Forecast\n", "Forecast USD portfolio NAV and risk from 100 positive float daily observations: equity index level, 10-year Treasury yield, and one EUR/USD FX rate. Instead of the unconditional forecast, ask the model a *what-if* question about the first future step: what does it expect for portfolio NAV and realized volatility **given** something about the other columns at that same step.\n", "\n", - "The `condition` query mode answers this in closed form from the model's joint distribution at one future position. A request names that position once, by its index into `query_times`, and attaches one condition per column it wants to fix:\n", + "The `condition` query mode answers this in closed form from the model's joint distribution at each future position. Every condition names the positions it covers by index into `query_times` (`query_time_indices`), and covers every position when it names none. `condition=` takes one condition or a list of them; conditions on different columns may cover the same position, while two conditions on the same column must cover disjoint positions:\n", "\n", "- An **equality condition** pins a column to a value (`EqualityCondition`). The pinned column leaves the read-out set, because reading it back would only repeat the request.\n", "- An **interval condition** confines a column to a range whose bounds may be open on either side (`IntervalCondition`). The column stays readable: what comes back is its distribution inside the range.\n", "\n", "The cell below asks for the same forecast twice, once without the block and once with it, because a conditional mean is only readable next to the unconditional one: what the what-if is worth is the shift between them, not the level of either.\n", "\n", - "Every column without a condition is a read-out column, and the response describes the conditional distribution of those columns at the conditioned position only, so `outputs.query_times` has exactly one entry however many `query_times` the request carried.\n", + "Every column without a condition is a read-out column. The response answers every entry of `query_times`, and a position no condition covers carries the unconditioned forecast. That is exact rather than a placeholder: the model draws each future position from its own joint, independently of the others, so a condition at one step says nothing about another. The comparison below shows it — the shift is zero at every step except the conditioned one.\n", "\n", "Every read-out the service serves reads that same conditional, and the sections below take it in turn: as a mean beside the unconditional one, as coherent joint draws, as quantiles inside a bounded band, and as the log density of values you already hold. The last two sections use the plausibility numbers to rank candidate what-ifs against each other.\n", "\n", @@ -78,7 +78,6 @@ "import pandas as pd\n", "\n", "from jointfm_client import (\n", - " ConditionBlock,\n", " EqualityCondition,\n", " JointFMClient,\n", " plan_forecast_columns,\n", @@ -94,9 +93,7 @@ "EQUITY_RALLY = 1.02\n", "EXPECTED_COLUMNS = FEATURE_COLUMNS + TARGET_COLUMNS\n", "QUERY_TIMES = list(range(INPUT_STEPS, INPUT_STEPS + OUTPUT_HORIZONS))\n", - "# Every column a request may read back: an equality condition removes its own\n", - "# column from the read-out set, and nothing else does.\n", - "READABLE_COLUMNS = [name for name in EXPECTED_COLUMNS if name != PINNED_COLUMN]\n", + "CONDITIONED_TIME = QUERY_TIMES[CONDITIONED_STEP]\n", "\n", "history = pd.read_csv(HISTORY_PATH, dtype=float)\n", "if list(history.columns) != EXPECTED_COLUMNS:\n", @@ -125,9 +122,10 @@ "\n", "last_equity_level = float(history[PINNED_COLUMN].iloc[-1])\n", "rallied_equity_level = last_equity_level * EQUITY_RALLY\n", - "rally = ConditionBlock(\n", - " query_time_index=CONDITIONED_STEP,\n", - " conditions=[EqualityCondition(column=PINNED_COLUMN, value=rallied_equity_level)],\n", + "rally = EqualityCondition(\n", + " column=PINNED_COLUMN,\n", + " value=rallied_equity_level,\n", + " query_time_indices=[CONDITIONED_STEP],\n", ")\n", "baseline = client.forecast_mean(\n", " history,\n", @@ -144,22 +142,18 @@ " seed=7,\n", " condition=rally,\n", ")\n", - "if result.query_times != (QUERY_TIMES[CONDITIONED_STEP],):\n", - " raise ValueError(\n", - " f\"Expected the conditioned position alone, got {result.query_times!r}\"\n", - " )\n", + "if result.query_times != tuple(QUERY_TIMES):\n", + " raise ValueError(f\"Expected every query time, got {result.query_times!r}\")\n", "if result.plausibility is None:\n", " raise ValueError(\"A condition response must carry its plausibility block\")\n", "print(\"log density of the pinned value:\", result.plausibility.equality_log_density)\n", "forecast = result.to_pandas_tidy()\n", - "expected_forecast_rows = len(plan.requested_columns)\n", + "expected_forecast_rows = len(plan.requested_columns) * len(QUERY_TIMES)\n", "if len(forecast) != expected_forecast_rows:\n", " raise ValueError(\n", " f\"Expected {expected_forecast_rows} forecast rows, got {len(forecast)}\"\n", " )\n", "\n", - "# The conditional answers one position, so joining on it keeps exactly the rows\n", - "# the two requests have in common.\n", "comparison = baseline.to_pandas_tidy().merge(\n", " forecast,\n", " on=[\"query_time\", \"requested_column\"],\n", @@ -168,6 +162,12 @@ "comparison[\"shift\"] = (\n", " comparison[\"value_given_the_rally\"] - comparison[\"value_unconditional\"]\n", ")\n", + "# Positions are drawn independently, so the pin can only move its own step.\n", + "moved_elsewhere = comparison[\n", + " (comparison[\"query_time\"] != CONDITIONED_TIME) & (comparison[\"shift\"] != 0.0)\n", + "]\n", + "if not moved_elsewhere.empty:\n", + " raise ValueError(f\"A step the rally does not cover moved:\\n{moved_elsewhere}\")\n", "comparison" ] }, @@ -188,6 +188,55 @@ { "cell_type": "markdown", "id": "5", + "metadata": { + "id": "condition-every-step-description", + "language": "markdown" + }, + "source": [ + "## Conditioning every step\n", + "A condition that names no `query_time_indices` covers every entry of `query_times`. Below, the equity index is held at the rallied level at all ten steps, and each step answers under that pin.\n", + "\n", + "Read it as ten what-ifs, one per step, not as a path along which the index stays put. Each step is conditioned on its own joint, and nothing links one step's joint to the next, so a pin at step 3 cannot inform step 4 — the same independence that left the uncovered steps unchanged above. What does add up across steps is the plausibility: `equality_log_density` is the sum of the ten per-step log densities, exactly, because the steps are independent. The same rule reaches `region_log_probability`, which sums over the steps an interval covers." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "6", + "metadata": { + "id": "condition-every-step", + "language": "python" + }, + "outputs": [], + "source": [ + "held_rally = EqualityCondition(column=PINNED_COLUMN, value=rallied_equity_level)\n", + "held = client.forecast_mean(\n", + " history,\n", + " query_times=QUERY_TIMES,\n", + " requested_columns=plan.requested_columns,\n", + " columns=plan.columns,\n", + " seed=7,\n", + " condition=held_rally,\n", + ")\n", + "if held.query_times != tuple(QUERY_TIMES):\n", + " raise ValueError(f\"Expected every query time, got {held.query_times!r}\")\n", + "if held.plausibility is None:\n", + " raise ValueError(\"A condition response must carry its plausibility block\")\n", + "print(\n", + " \"log density of the held level over every step:\",\n", + " held.plausibility.equality_log_density,\n", + ")\n", + "held_path = baseline.to_pandas_wide().merge(\n", + " held.to_pandas_wide(),\n", + " on=\"query_time\",\n", + " suffixes=(\"_unconditional\", \"_held\"),\n", + ")\n", + "held_path" + ] + }, + { + "cell_type": "markdown", + "id": "7", "metadata": { "id": "condition-draws-description", "language": "markdown" @@ -196,13 +245,13 @@ "## Scenarios, not summaries\n", "A mean answers *where*, and a band answers *how wide*, but a portfolio question usually needs whole futures: draws. Every row below is one coherent joint scenario, because the service draws the read-out columns together from the same conditional mixture — the NAV and the volatility in one row belong to each other, so a function of several columns can be evaluated row by row. A quantile table cannot answer that: the 90th percentile of NAV and the 90th percentile of volatility need not describe any single future.\n", "\n", - "`diagnostics.condition_draws` reports how many draws stand behind the answer. Oversized sample requests are split into batches against the deployment's advertised cap and merged locally, and the merged count is what this field reports." + "The table shows the conditioned step; the other steps come back too, drawn unconditionally. `diagnostics.condition_draws` reports how many draws stand behind the answer. Oversized sample requests are split into batches against the deployment's advertised cap and merged locally, and the merged count is what this field reports." ] }, { "cell_type": "code", "execution_count": null, - "id": "6", + "id": "8", "metadata": { "id": "condition-draws", "language": "python" @@ -220,36 +269,35 @@ " seed=7,\n", " condition=rally,\n", ")\n", - "if scenarios.query_times != (QUERY_TIMES[CONDITIONED_STEP],):\n", - " raise ValueError(\n", - " f\"Expected the conditioned position alone, got {scenarios.query_times!r}\"\n", - " )\n", + "if scenarios.query_times != tuple(QUERY_TIMES):\n", + " raise ValueError(f\"Expected every query time, got {scenarios.query_times!r}\")\n", "print(\"draws behind the answer:\", scenarios.diagnostics.condition_draws)\n", - "scenarios.to_pandas_wide()" + "draws = scenarios.to_pandas_wide()\n", + "draws[draws[\"query_time\"] == CONDITIONED_TIME].reset_index(drop=True)" ] }, { "cell_type": "markdown", - "id": "7", + "id": "9", "metadata": { "id": "condition-interval-description", "language": "markdown" }, "source": [ "## Interval condition\n", - "Now confine the 10-year yield to a band around its last observed value instead of pinning it, and pin the equity index at the same time: a request may mix both kinds across the columns of one position. The response then carries `region_log_probability`, and the yield column itself stays readable because its distribution inside the band is a genuine answer.\n", + "Now confine the 10-year yield to a band around its last observed value instead of pinning it, and pin the equity index at the same time: conditions on different columns may cover the same step, which is how one step mixes both kinds. The response then carries `region_log_probability`, and the yield column itself stays readable because its distribution inside the band is a genuine answer.\n", "\n", "That number is the log probability of the band **given the pinned equity index**, not the band's own probability, because the two condition kinds compose in a fixed order and the band is measured on the distribution the pin has already reduced. The two numbers therefore chain rather than describe separate things, and adding them gives the plausibility of the whole request — `log(density(pin) * P(band | pin))` — with no independence assumed anywhere. A request carrying only interval conditions has nothing to combine, and its `region_log_probability` alone is the joint probability of everything it asked about.\n", "\n", "With one interval column, as here, the region probability is exact — one difference of distribution functions per mixture component — and `diagnostics.interval_estimator` stays empty, because there is no estimate to characterize. Bounding a second column makes the region a box with no closed form, which the service estimates numerically and then reports the accounting for — the next section does exactly that.\n", "\n", - "A band is also where the quantile read-out earns its place over the mean: what the banded column comes back with is a distribution inside its own range, so the cell asserts every quantile of it lands there." + "A band is also where the quantile read-out earns its place over the mean: what the banded column comes back with is a distribution inside its own range, so the cell asserts every quantile of it lands there at the step the band covers." ] }, { "cell_type": "code", "execution_count": null, - "id": "8", + "id": "10", "metadata": { "id": "condition-interval-example", "language": "python" @@ -268,15 +316,12 @@ "yield_lower = last_yield - YIELD_BAND_HALF_WIDTH\n", "yield_upper = last_yield + YIELD_BAND_HALF_WIDTH\n", "yield_band = IntervalCondition(\n", - " column=\"treasury_10y_yield\", lower=yield_lower, upper=yield_upper\n", - ")\n", - "rally_with_yield_band = ConditionBlock(\n", - " query_time_index=CONDITIONED_STEP,\n", - " conditions=[\n", - " EqualityCondition(column=PINNED_COLUMN, value=rallied_equity_level),\n", - " yield_band,\n", - " ],\n", + " column=\"treasury_10y_yield\",\n", + " lower=yield_lower,\n", + " upper=yield_upper,\n", + " query_time_indices=[CONDITIONED_STEP],\n", ")\n", + "rally_with_yield_band = [rally, yield_band]\n", "banded = client.forecast_quantiles(\n", " history,\n", " query_times=QUERY_TIMES,\n", @@ -292,7 +337,10 @@ "print(\"log probability of the yield band:\", banded.plausibility.region_log_probability)\n", "if banded.diagnostics.interval_estimator is not None:\n", " raise ValueError(\"A single interval column is exact and reports no estimator\")\n", - "band = banded.to_pandas_wide()\n", + "quantile_frame = banded.to_pandas_wide()\n", + "band = quantile_frame[quantile_frame[\"query_time\"] == CONDITIONED_TIME].reset_index(\n", + " drop=True\n", + ")\n", "outside_band = band[\n", " (band[\"treasury_10y_yield\"] < yield_lower - BAND_TOLERANCE)\n", " | (band[\"treasury_10y_yield\"] > yield_upper + BAND_TOLERANCE)\n", @@ -304,7 +352,7 @@ }, { "cell_type": "markdown", - "id": "9", + "id": "11", "metadata": { "id": "condition-box-description", "language": "markdown" @@ -313,13 +361,13 @@ "## When the region has to be estimated\n", "Bounding the FX rate as well turns the region into a box. That has no closed form, so the service estimates its probability with a quasi-random walk and only then reports the accounting in `diagnostics.interval_estimator`: `points` is how many quasi-random points went into the estimate and `effective_sample_size` is Kish's effective sample size of the weights behind them. A small effective sample size against a large point count means a deep-tail box where few points carry the answer. The cell asserts both sides of that rule — absent for the single band above, present here.\n", "\n", - "Drawing from a boxed conditional also shows what the box does to the draws. Each row pairs one truncated draw of the bounded columns with a read-out drawn from *that* draw's conditional, so the rows are exact draws from the joint conditional rather than separately summarized margins, and every one of them lands inside both bands." + "Drawing from a boxed conditional also shows what the box does to the draws. Each row pairs one truncated draw of the bounded columns with a read-out drawn from *that* draw's conditional, so the rows are exact draws from the joint conditional rather than separately summarized margins, and every one of them lands inside both bands at the step they cover." ] }, { "cell_type": "code", "execution_count": null, - "id": "10", + "id": "12", "metadata": { "id": "condition-box", "language": "python" @@ -331,15 +379,13 @@ "last_fx_rate = float(history[\"eur_usd_rate\"].iloc[-1])\n", "fx_lower = last_fx_rate - FX_BAND_HALF_WIDTH\n", "fx_upper = last_fx_rate + FX_BAND_HALF_WIDTH\n", - "fx_band = IntervalCondition(column=\"eur_usd_rate\", lower=fx_lower, upper=fx_upper)\n", - "rally_with_two_bands = ConditionBlock(\n", - " query_time_index=CONDITIONED_STEP,\n", - " conditions=[\n", - " EqualityCondition(column=PINNED_COLUMN, value=rallied_equity_level),\n", - " yield_band,\n", - " fx_band,\n", - " ],\n", + "fx_band = IntervalCondition(\n", + " column=\"eur_usd_rate\",\n", + " lower=fx_lower,\n", + " upper=fx_upper,\n", + " query_time_indices=[CONDITIONED_STEP],\n", ")\n", + "rally_with_two_bands = [rally, yield_band, fx_band]\n", "boxed = client.forecast_samples(\n", " history,\n", " query_times=QUERY_TIMES,\n", @@ -362,7 +408,10 @@ ")\n", "print(\"estimator points:\", estimator.points)\n", "print(\"effective sample size:\", estimator.effective_sample_size)\n", - "box_draws = boxed.to_pandas_wide()\n", + "boxed_frame = boxed.to_pandas_wide()\n", + "box_draws = boxed_frame[boxed_frame[\"query_time\"] == CONDITIONED_TIME].reset_index(\n", + " drop=True\n", + ")\n", "outside_box = box_draws[\n", " (box_draws[\"treasury_10y_yield\"] < yield_lower - BAND_TOLERANCE)\n", " | (box_draws[\"treasury_10y_yield\"] > yield_upper + BAND_TOLERANCE)\n", @@ -376,7 +425,7 @@ }, { "cell_type": "markdown", - "id": "11", + "id": "13", "metadata": { "id": "condition-ranking-description", "language": "markdown" @@ -391,7 +440,7 @@ { "cell_type": "code", "execution_count": null, - "id": "12", + "id": "14", "metadata": { "id": "condition-ranking", "language": "python" @@ -409,14 +458,18 @@ " requested_columns=plan.requested_columns,\n", " columns=plan.columns,\n", " seed=7,\n", - " condition=ConditionBlock(\n", - " query_time_index=CONDITIONED_STEP,\n", - " conditions=[EqualityCondition(column=PINNED_COLUMN, value=candidate_level)],\n", + " condition=EqualityCondition(\n", + " column=PINNED_COLUMN,\n", + " value=candidate_level,\n", + " query_time_indices=[CONDITIONED_STEP],\n", " ),\n", " )\n", " if candidate_result.plausibility is None:\n", " raise ValueError(\"A condition response must carry its plausibility block\")\n", - " conditional_mean = candidate_result.to_pandas_wide().iloc[0]\n", + " candidate_frame = candidate_result.to_pandas_wide()\n", + " conditional_mean = candidate_frame[\n", + " candidate_frame[\"query_time\"] == CONDITIONED_TIME\n", + " ].iloc[0]\n", " rankings.append(\n", " {\n", " \"rally\": candidate,\n", @@ -435,7 +488,7 @@ }, { "cell_type": "markdown", - "id": "13", + "id": "15", "metadata": { "id": "condition-score-description", "language": "markdown" @@ -444,7 +497,7 @@ "## Scoring an outcome you already have\n", "Every read-out so far answers *what does the model expect*. `log_prob` asks the opposite — how plausible are values I already hold — and is the only return mode where the caller supplies the future instead of receiving it. Under a condition it scores those values against the conditional, which is how two candidate outcomes of the same what-if are compared.\n", "\n", - "Three rules follow from scoring a joint rather than a projection of one, and the client enforces all three before anything is sent. `query_rows` carries one observed row per entry of `query_times`, each with a value for every declared column. `requested_columns` must name every declared column in declared order — a narrower projection would score a different distribution than the caller means to ask about — and a condition excuses nothing, so omitting the field is the natural spelling. The deployment refuses a row that contradicts the condition, a pinned column given another value or a banded column outside its range, instead of scoring it.\n", + "Three rules follow from scoring a joint rather than a projection of one, and the client enforces all three before anything is sent. `query_rows` carries one observed row per entry of `query_times`, each with a value for every declared column. `requested_columns` must name every declared column in declared order — a narrower projection would score a different distribution than the caller means to ask about — and a condition excuses nothing, so omitting the field is the natural spelling. The deployment refuses a row that contradicts a condition covering its step, a pinned column given another value or a banded column outside its range, instead of scoring it.\n", "\n", "The scores below are log densities of whole rows, so the unit caveat from the plausibility section applies unchanged: read their difference, not their level." ] @@ -452,7 +505,7 @@ { "cell_type": "code", "execution_count": null, - "id": "14", + "id": "16", "metadata": { "id": "condition-score", "language": "python" @@ -481,7 +534,6 @@ " history,\n", " query_times=SCORED_QUERY_TIMES,\n", " query_rows=pd.DataFrame([outcome], columns=EXPECTED_COLUMNS),\n", - " requested_columns=READABLE_COLUMNS,\n", " columns=plan.columns,\n", " seed=7,\n", " condition=rally,\n", diff --git a/src/jointfm_client/__init__.py b/src/jointfm_client/__init__.py index ffe9a94..77a9900 100644 --- a/src/jointfm_client/__init__.py +++ b/src/jointfm_client/__init__.py @@ -69,13 +69,14 @@ PACKAGE_VERSION, PREDICT_REQUEST_TYPE, SCHEMA_VERSION, - ConditionBlock, + Condition, ConditionKind, ConditionPlausibility, EqualityCondition, IntervalCondition, IntervalEstimator, require_condition_support, + resolve_conditions, MeanForecastResult, QuantileForecast, QuantileForecastResult, @@ -224,13 +225,14 @@ "PREDICT_REQUEST_TYPE", "RetryConfig", "SCHEMA_VERSION", - "ConditionBlock", + "Condition", "ConditionKind", "ConditionPlausibility", "EqualityCondition", "IntervalCondition", "IntervalEstimator", "require_condition_support", + "resolve_conditions", "SampleForecastResult", "SUPPORTED_COLUMN_MODALITIES", "SUPPORTED_COLUMN_ROLES", diff --git a/src/jointfm_client/adapters.py b/src/jointfm_client/adapters.py index d1aa57d..8872431 100644 --- a/src/jointfm_client/adapters.py +++ b/src/jointfm_client/adapters.py @@ -24,7 +24,7 @@ from jointfm_client.configuration import DEFAULT_FORECAST_SCHEMA_VERSION from jointfm_client.contract import ( DEFAULT_CALENDAR_ID, - ConditionBlock, + Condition, ColumnModality, ColumnRole, ColumnSpec, @@ -304,14 +304,15 @@ def build_forecast_payload_from_dataframe( time_value_columns: Sequence[str] | Mapping[str, TimeValueKind] | None = None, nullable_columns: Sequence[str] | None = None, bounds: ColumnBounds | None = None, - condition: ConditionBlock | None = None, + condition: Condition | Sequence[Condition] | None = None, query_rows: Any | None = None, ) -> dict[str, Any]: """Build a validated forecast payload from a pandas ``DataFrame``. - Passing ``condition`` makes this a conditioning request: the payload then - carries ``query_mode='condition'`` and the block, and the deployment - answers the conditional at the one future position the block names. + Passing ``condition`` — one condition or a list of them — makes this a + conditioning request: the payload then carries ``query_mode='condition'`` + and the conditions, and the deployment answers every future position under + the conditions covering it. ``query_rows`` carries the *observed* values at ``query_times`` that ``return_mode='log_prob'`` scores — a ``DataFrame`` shaped like ``frame``, diff --git a/src/jointfm_client/client.py b/src/jointfm_client/client.py index 8a6a5bd..7c08bd7 100644 --- a/src/jointfm_client/client.py +++ b/src/jointfm_client/client.py @@ -36,7 +36,7 @@ load_configuration, ) from jointfm_client.contract import ( - ConditionBlock, + Condition, DEFAULT_CALENDAR_ID, HEALTH_REQUEST_TYPE, SCHEMA_VERSION, @@ -313,15 +313,21 @@ def forecast( nullable_columns: Sequence[str] | None = None, bounds: Mapping[str, tuple[float | int | None, float | int | None]] | None = None, - condition: ConditionBlock | None = None, + condition: Condition | Sequence[Condition] | None = None, query_rows: Any | None = None, ) -> ForecastResponse: """Build and submit a forecast request from tabular history inputs. - Passing ``condition`` asks the deployment for the conditional at the one - future position the block names, instead of the unconditional forecast. - The deployment's advertised capability is checked first, so a deployment - that cannot condition is refused here rather than after a round trip. + Passing ``condition`` — one ``EqualityCondition`` or + ``IntervalCondition``, or a list of them — asks the deployment for the + forecast given those conditions. Each condition covers the positions its + ``query_time_indices`` names, every position when that is ``None``, and + two conditions on one column must cover disjoint positions. The answer + still covers every entry of ``query_times``: a position no condition + covers carries the unconditioned forecast, which is exact because the + deployment draws each position independently. The deployment's + advertised capability is checked first, so a deployment that cannot + condition is refused here rather than after a round trip. ``query_rows`` carries the observed values at ``query_times`` that ``return_mode='log_prob'`` scores, in the same shape as ``history``. @@ -417,12 +423,12 @@ def forecast_mean( requested_columns: Sequence[str | int] | None = None, model_version: str | None = None, seed: int | None = None, - condition: ConditionBlock | None = None, + condition: Condition | Sequence[Condition] | None = None, ) -> MeanForecastResult: """Forecast mean values through the shared forecast validation path. - With ``condition`` the mean is the conditional mean at the one future - position the block names; see :meth:`forecast`. + With ``condition`` the mean at each covered position is the conditional + mean there; see :meth:`forecast`. """ return cast( MeanForecastResult, @@ -454,12 +460,12 @@ def forecast_samples( model_version: str | None = None, n_samples: int | None = None, seed: int | None = None, - condition: ConditionBlock | None = None, + condition: Condition | Sequence[Condition] | None = None, ) -> SampleForecastResult: """Forecast sample paths through the shared forecast validation path. - With ``condition`` the draws come from the conditional at the one future - position the block names; see :meth:`forecast`. + With ``condition`` the draws at each covered position come from the + conditional there; see :meth:`forecast`. """ return cast( SampleForecastResult, @@ -493,12 +499,12 @@ def forecast_quantiles( n_samples: int | None = None, quantiles: Sequence[float | int] | None = None, seed: int | None = None, - condition: ConditionBlock | None = None, + condition: Condition | Sequence[Condition] | None = None, ) -> QuantileForecastResult: """Forecast quantiles through the shared forecast validation path. - With ``condition`` the quantiles describe the conditional at the one - future position the block names; see :meth:`forecast`. + With ``condition`` the quantiles at each covered position describe the + conditional there; see :meth:`forecast`. """ return cast( QuantileForecastResult, @@ -532,7 +538,7 @@ def forecast_log_prob( requested_columns: Sequence[str | int] | None = None, model_version: str | None = None, seed: int | None = None, - condition: ConditionBlock | None = None, + condition: Condition | Sequence[Condition] | None = None, ) -> LogProbResult: """Score observed future values through the shared forecast validation path. @@ -544,9 +550,9 @@ def forecast_log_prob( declared column in declared order, which is what omitting it already resolves to. - With ``condition`` the score is taken under the conditional at the one - future position the block names, and the service refuses rows that - contradict the condition instead of scoring them; see :meth:`forecast`. + With ``condition`` each covered row is scored under the conditional at + its position, and the service refuses a row that contradicts a + condition covering it instead of scoring it; see :meth:`forecast`. """ return cast( LogProbResult, @@ -721,7 +727,7 @@ def _forecast_payload_from_rows( use_local_normalized_time: bool, calendar_id: str, timezone: str | None, - condition: ConditionBlock | None = None, + condition: Condition | Sequence[Condition] | None = None, query_rows: Any | None = None, ) -> dict[str, Any]: if schema is None: diff --git a/src/jointfm_client/contract.py b/src/jointfm_client/contract.py index 2b028b4..598875e 100644 --- a/src/jointfm_client/contract.py +++ b/src/jointfm_client/contract.py @@ -37,7 +37,7 @@ # pyproject.toml the way a hand-maintained literal here did. PACKAGE_VERSION: Final = importlib.metadata.version(DISTRIBUTION_NAME) -SCHEMA_VERSION: Final = "v4" +SCHEMA_VERSION: Final = "v5" # The service sums and averages its log densities in the model's own precision, # so its summaries differ from a recomputation in the last few digits. DERIVED_SCORE_TOLERANCE: Final = 1e-6 @@ -311,19 +311,37 @@ def to_payload(self) -> dict[str, Any]: return payload +def _require_query_time_indices(value: Any) -> tuple[int, ...] | None: + """Validate the positions one condition covers, ``None`` meaning every one.""" + if value is None: + return None + indices = _require_sequence(value, field="condition.query_time_indices") + for index in indices: + if isinstance(index, bool) or not isinstance(index, int): + raise ValueError("condition.query_time_indices must hold integers") + if index < 0: + raise ValueError("condition.query_time_indices must not be negative") + if len(set(indices)) != len(indices): + raise ValueError("condition.query_time_indices names a position more than once") + return tuple(indices) + + @dataclass(frozen=True, slots=True) class EqualityCondition: - """One column of the conditioned position pinned to a value. - - An equality condition is an event of probability zero under a continuous - head, so the deployment answers it analytically rather than by filtering - draws. It fixes the column's value, so a projection that names the column - reads that value back rather than a model output — which is what lets a - scenario answer line up column for column with an unconditioned one. + """One column pinned to a value at the future positions it covers. + + ``query_time_indices`` names those positions by index into the request's + ``query_times``; ``None`` covers every one of them. An equality condition + is an event of probability zero under a continuous head, so the deployment + answers it analytically rather than by filtering draws. It fixes the + column's value, so a projection that names the column reads that value back + at every covered position rather than a model output — which is what lets + a scenario answer line up column for column with an unconditioned one. """ column: str value: float + query_time_indices: Sequence[int] | None = None def __post_init__(self) -> None: """Reject a pin the service would refuse anyway.""" @@ -332,25 +350,37 @@ def __post_init__(self) -> None: raise ValueError("condition value must be a number") if not math.isfinite(float(self.value)): raise ValueError("condition value must be finite") + object.__setattr__( + self, + "query_time_indices", + _require_query_time_indices(self.query_time_indices), + ) def to_payload(self) -> dict[str, Any]: """Return this condition as JSON-compatible payload fields.""" - return {"column": self.column, "kind": "equality", "value": float(self.value)} + return { + "column": self.column, + "kind": "equality", + "value": float(self.value), + "query_time_indices": _query_time_indices_payload(self), + } @dataclass(frozen=True, slots=True) class IntervalCondition: - """One column of the conditioned position confined to a range. + """One column confined to a range at the future positions it covers. - ``None`` on either side leaves it open. The region carries positive mass, - so the deployment answers it by composition on the reduced mixture, and the - column stays readable: what it returns is its distribution inside the - region. + ``None`` on either side leaves the range open, and ``query_time_indices`` + of ``None`` covers every requested position. The region carries positive + mass, so the deployment answers it by composition on the reduced mixture, + and the column stays readable: what it returns is its distribution inside + the region. """ column: str lower: float | None = None upper: float | None = None + query_time_indices: Sequence[int] | None = None def __post_init__(self) -> None: """Reject a range that bounds nothing or bounds it backwards.""" @@ -374,6 +404,11 @@ def __post_init__(self) -> None: raise ValueError( f"condition needs lower < upper, got [{self.lower}, {self.upper}]" ) + object.__setattr__( + self, + "query_time_indices", + _require_query_time_indices(self.query_time_indices), + ) def to_payload(self) -> dict[str, Any]: """Return this condition as JSON-compatible payload fields.""" @@ -382,78 +417,85 @@ def to_payload(self) -> dict[str, Any]: "kind": "interval", "lower": None if self.lower is None else float(self.lower), "upper": None if self.upper is None else float(self.upper), + "query_time_indices": _query_time_indices_payload(self), } -@dataclass(frozen=True, slots=True) -class ConditionBlock: - """Every condition of one request, against one future position. +Condition: TypeAlias = EqualityCondition | IntervalCondition - The position is named once, by its index into the request's - ``query_times``, because conditioning relates variables to each other - within a position and never across horizons. A statement spanning - horizons is out of reach of this mode and cannot be expressed here. - """ - query_time_index: int - conditions: Sequence[EqualityCondition | IntervalCondition] +def _query_time_indices_payload(condition: Condition) -> list[int] | None: + """Serialize a condition's positions, ``None`` staying the wire's ``null``.""" + if condition.query_time_indices is None: + return None + return list(condition.query_time_indices) - def __post_init__(self) -> None: - """Validate the block independently of any schema or deployment.""" - if isinstance(self.query_time_index, bool) or not isinstance( - self.query_time_index, int - ): - raise ValueError("query_time_index must be an integer") - if self.query_time_index < 0: - raise ValueError("query_time_index must not be negative") - conditions = _require_sequence(self.conditions, field="conditions") - if not conditions: - raise ValueError("a condition block must carry at least one condition") - seen: set[str] = set() - for index, condition in enumerate(conditions): - if not isinstance(condition, (EqualityCondition, IntervalCondition)): - raise ValueError( - f"conditions[{index}] must be an EqualityCondition or an " - "IntervalCondition" - ) - if condition.column in seen: - raise ValueError( - f"column {condition.column!r} carries more than one condition" - ) - seen.add(condition.column) - @property - def pinned_columns(self) -> tuple[str, ...]: - """Columns an equality condition fixes, whose answer is the request's own value.""" - return tuple( - condition.column - for condition in self.conditions - if isinstance(condition, EqualityCondition) - ) +def _covers(condition: Condition, position: int) -> bool: + """Return whether ``condition`` applies at future ``position``.""" + return ( + condition.query_time_indices is None or position in condition.query_time_indices + ) - @property - def conditioned_columns(self) -> tuple[str, ...]: - """Every column this block conditions, of either kind.""" - return tuple(condition.column for condition in self.conditions) - @property - def kinds(self) -> tuple[ConditionKind, ...]: - """The condition kinds this block uses, which a deployment must advertise.""" - kinds: list[ConditionKind] = [] - for condition in self.conditions: - kind: ConditionKind = ( - "equality" if isinstance(condition, EqualityCondition) else "interval" +def resolve_conditions( + condition: Condition | Sequence[Condition], +) -> tuple[Condition, ...]: + """Normalize one condition or a list of them into a validated tuple. + + Conditions on different columns may cover the same positions, which is how + one position mixes a pin and a band. Two conditions on the *same* column + must cover disjoint positions, because a column carries at most one + condition per position; ``query_time_indices=None`` covers every position, + so it overlaps any other condition on its column. + + Raises: + ValueError: when the list is empty, holds something other than an + ``EqualityCondition`` or ``IntervalCondition``, or puts two + conditions on one column at a shared position. + """ + conditions: tuple[Any, ...] = ( + (condition,) + if isinstance(condition, (EqualityCondition, IntervalCondition)) + else tuple(_require_sequence(condition, field="condition")) + ) + for index, entry in enumerate(conditions): + if not isinstance(entry, (EqualityCondition, IntervalCondition)): + raise ValueError( + f"condition[{index}] must be an EqualityCondition or an " + "IntervalCondition" ) - if kind not in kinds: - kinds.append(kind) - return tuple(kinds) + for later_index, later in enumerate(conditions): + for earlier_index, earlier in enumerate(conditions[:later_index]): + if earlier.column != later.column: + continue + if earlier.query_time_indices is None or later.query_time_indices is None: + overlap = "every position" + else: + shared = sorted( + set(earlier.query_time_indices) & set(later.query_time_indices) + ) + if not shared: + continue + overlap = f"query_time_indices {shared}" + raise ValueError( + f"column {later.column!r} carries more than one condition at " + f"{overlap}: condition[{later_index}] overlaps " + f"condition[{earlier_index}]" + ) + return cast(tuple[Condition, ...], conditions) - def to_payload(self) -> dict[str, Any]: - """Return the block as JSON-compatible payload fields.""" - return { - "query_time_index": self.query_time_index, - "conditions": [condition.to_payload() for condition in self.conditions], - } + +def condition_kinds(conditions: Sequence[Condition]) -> tuple[ConditionKind, ...]: + """Return the condition kinds ``conditions`` use, which a deployment must advertise.""" + kinds: list[ConditionKind] = [] + for condition in conditions: + kind: ConditionKind = ( + "equality" if isinstance(condition, EqualityCondition) else "interval" + ) + if kind not in kinds: + kinds.append(kind) + return tuple(kinds) @dataclass(frozen=True, slots=True) @@ -506,7 +548,7 @@ class ForecastRequest: quantiles: Sequence[float | int] | None = None seed: int | None = None query_row_ids: Sequence[int] | None = None - condition: ConditionBlock | None = None + condition: Condition | Sequence[Condition] | None = None query_rows: Sequence[Mapping[str, Any]] | None = None def __post_init__(self) -> None: @@ -531,7 +573,7 @@ def __post_init__(self) -> None: if (self.metadata.query_mode == "condition") != (self.condition is not None): raise ValueError( - "query_mode='condition' and a condition block go together: one " + "query_mode='condition' and a condition go together: one " "without the other means a different request than the caller wrote" ) if self.condition is not None: @@ -553,34 +595,50 @@ def __post_init__(self) -> None: self._validate_query_rows() def _validate_condition(self) -> None: - """Check the condition block against this request's schema and horizons. + """Check the conditions against this request's schema and positions. The deployment rejects all of this too, before executing anything; the - point of repeating it here is that a caller learns which column name it - got wrong without paying for a round trip. + point of repeating it here is that a caller learns which column name or + position it got wrong without paying for a round trip. """ - block = self.condition - assert block is not None + assert self.condition is not None + conditions = resolve_conditions(self.condition) declared = {column.name for column in self.schema.columns} - unknown = [name for name in block.conditioned_columns if name not in declared] + unknown = [ + condition.column + for condition in conditions + if condition.column not in declared + ] if unknown: raise ValueError(f"condition references undeclared columns {unknown}") horizon_count = len(_require_sequence(self.query_times, field="query_times")) - if block.query_time_index >= horizon_count: - raise ValueError( - f"condition.query_time_index {block.query_time_index} is outside the " - f"{horizon_count} requested future positions" - ) - if len(set(block.conditioned_columns)) >= len(declared): - if len(block.pinned_columns) == len(set(block.conditioned_columns)): + for condition in conditions: + outside = [ + index + for index in condition.query_time_indices or () + if index >= horizon_count + ] + if outside: raise ValueError( - "pinning every column leaves nothing to read out; the joint " - "density of those values is what return_mode='log_prob' answers" + f"condition on {condition.column!r} names query_time_indices " + f"{outside} outside the {horizon_count} requested future positions" + ) + for position in range(horizon_count): + covering = [ + condition for condition in conditions if _covers(condition, position) + ] + if len(covering) < len(declared): + continue + if all(isinstance(condition, EqualityCondition) for condition in covering): + raise ValueError( + f"pinning every column at query_time_indices position {position} " + "leaves nothing to read out; the joint density of those values " + "is what return_mode='log_prob' answers" ) raise ValueError( - "a request must leave at least one column unconditioned: what it " - "reads out is the conditional distribution of the columns it does " - "not condition" + "a request must leave at least one column unconditioned at " + f"query_time_indices position {position}: what it reads out is " + "the conditional distribution of the columns it does not condition" ) def _validate_query_rows(self) -> None: @@ -590,7 +648,7 @@ def _validate_query_rows(self) -> None: must carry every declared column and ``requested_columns`` must name every one of them in declared order: a narrower or reordered projection would score a different distribution than the caller believes it asked - about. A ``condition`` block changes nothing here — its columns are + about. A ``condition`` changes nothing here — its columns are declared columns like any other — and omitting the field already resolves to exactly this projection. @@ -663,7 +721,10 @@ def to_payload(self) -> dict[str, Any]: if self.seed is not None: payload["seed"] = self.seed if self.condition is not None: - payload["condition"] = self.condition.to_payload() + payload["conditions"] = [ + condition.to_payload() + for condition in resolve_conditions(self.condition) + ] if self.query_rows is not None: payload["query_rows"] = _serialize_rows(self.query_rows, field="query_rows") return payload @@ -850,9 +911,12 @@ def from_payload(cls, payload: Mapping[str, Any]) -> Self: class IntervalEstimator: """How a multi-column region's probability was estimated. - Present only when the interval block spans more than one column, which is - the only case where a component's box probability has to be estimated; one + Present only when some position bounds more than one column, which is the + only case where a component's box probability has to be estimated; one interval column is exact, and a number that never varies would say nothing. + When several positions are estimated, the accounting describes the one with + the smallest effective sample size, because that position bounds how far + the whole response can be trusted. """ points: int @@ -883,14 +947,17 @@ class ConditionPlausibility: kind. The service reports them and never refuses on them, so acting on them is the caller's decision. - Each number is joint over its whole condition block rather than one value + Each number is joint over every condition of its kind rather than one value per column: ``equality_log_density`` covers every pinned value at once, and ``region_log_probability`` covers the whole box at once. Do not rebuild either by multiplying per-column numbers, which would assume the columns are - independent and so discard what a joint model is for. + independent and so discard what a joint model is for. Across positions the + numbers *are* sums: the deployment draws each future position from its own + joint, independently of the others, so the value for a request covering + several positions is the sum of the per-position values, exactly. ``region_log_probability`` is measured *after* the pinned columns are - applied, so on a request carrying both kinds it is the log probability of + applied, so at a position carrying both kinds it is the log probability of the box **given the pins**, not the box's own. The two therefore chain, and their sum is the plausibility of the whole condition set:: @@ -1628,7 +1695,7 @@ def build_forecast_payload( seed: int | None = None, schema_version: str = SCHEMA_VERSION, query_mode: QueryMode = "forecast", - condition: ConditionBlock | None = None, + condition: Condition | Sequence[Condition] | None = None, query_rows: Sequence[Mapping[str, Any]] | None = None, ) -> dict[str, Any]: """Build a validated JSON-compatible forecast request payload.""" @@ -1653,7 +1720,7 @@ def build_forecast_payload( def require_condition_support( metadata: HealthMetadata, - block: ConditionBlock, + condition: Condition | Sequence[Condition], ) -> None: """Refuse a conditioning request the mounted deployment has not advertised. @@ -1664,7 +1731,7 @@ def require_condition_support( Raises: UnsupportedServiceContractError: when the deployment serves no - ``condition`` mode, or none of the condition kinds this block uses. + ``condition`` mode, or not every condition kind ``condition`` uses. """ if "condition" not in metadata.supported_query_modes: raise UnsupportedServiceContractError( @@ -1672,7 +1739,9 @@ def require_condition_support( f"queries; it advertises {list(metadata.supported_query_modes)}" ) missing = [ - kind for kind in block.kinds if kind not in metadata.supported_condition_kinds + kind + for kind in condition_kinds(resolve_conditions(condition)) + if kind not in metadata.supported_condition_kinds ] if missing: raise UnsupportedServiceContractError( @@ -2064,18 +2133,6 @@ def _forecast_response_expectations( ): requested_columns = tuple(requested_column_values) - condition_value = request_payload.get("condition") - if condition_value is not None and query_times is not None: - condition_block = _require_mapping( - condition_value, field="request_payload.condition" - ) - conditioned_index = condition_block.get("query_time_index") - if isinstance(conditioned_index, int) and not isinstance( - conditioned_index, bool - ): - if 0 <= conditioned_index < len(query_times): - query_times = (query_times[conditioned_index],) - query_mode_value = request_payload.get("query_mode") return_mode_value = request_payload.get("return_mode") quantiles_value = request_payload.get("quantiles") diff --git a/tests/fixtures/condition_interval_response.json b/tests/fixtures/condition_interval_response.json index f217332..d076318 100644 --- a/tests/fixtures/condition_interval_response.json +++ b/tests/fixtures/condition_interval_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", @@ -7,9 +7,10 @@ "query_mode": "condition", "return_mode": "mean", "outputs": { - "query_times": [3], + "query_times": [2, 3], "requested_columns": ["target"], "mean": [ + [12.0], [14.25] ], "samples": null, @@ -21,7 +22,7 @@ }, "diagnostics": { "history_rows": 2, - "horizon_count": 1, + "horizon_count": 2, "seed": 7, "condition_draws": 500, "interval_estimator": { diff --git a/tests/fixtures/condition_log_prob_request.json b/tests/fixtures/condition_log_prob_request.json index b6ef705..c3e5c74 100644 --- a/tests/fixtures/condition_log_prob_request.json +++ b/tests/fixtures/condition_log_prob_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "condition", "return_mode": "log_prob", @@ -27,16 +27,14 @@ ], "query_times": [2, 3], "requested_columns": ["driver", "target"], - "condition": { - "query_time_index": 1, - "conditions": [ - { - "column": "driver", - "kind": "equality", - "value": 1.5 - } - ] - }, + "conditions": [ + { + "column": "driver", + "kind": "equality", + "value": 1.5, + "query_time_indices": [1] + } + ], "query_rows": [ { "driver": 1.2, diff --git a/tests/fixtures/condition_log_prob_response.json b/tests/fixtures/condition_log_prob_response.json index 04513ea..a78c15c 100644 --- a/tests/fixtures/condition_log_prob_response.json +++ b/tests/fixtures/condition_log_prob_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", @@ -7,18 +7,18 @@ "query_mode": "condition", "return_mode": "log_prob", "outputs": { - "query_times": [3], + "query_times": [2, 3], "requested_columns": ["driver", "target"], "mean": null, "samples": null, "quantiles": null, "log_prob": { - "values": [-1.75], - "nll_values": [1.75], - "total": -1.75, - "mean": -1.75, - "nll_total": 1.75, - "nll_mean": 1.75 + "values": [-2.5, -1.75], + "nll_values": [2.5, 1.75], + "total": -4.25, + "mean": -2.125, + "nll_total": 4.25, + "nll_mean": 2.125 } }, "plausibility": { @@ -27,7 +27,7 @@ }, "diagnostics": { "history_rows": 2, - "horizon_count": 1, + "horizon_count": 2, "seed": 7, "condition_draws": 500, "interval_estimator": null diff --git a/tests/fixtures/condition_mean_request.json b/tests/fixtures/condition_mean_request.json index f506b20..ecd9608 100644 --- a/tests/fixtures/condition_mean_request.json +++ b/tests/fixtures/condition_mean_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "condition", "return_mode": "mean", @@ -27,15 +27,13 @@ ], "query_times": [2, 3], "requested_columns": ["target"], - "condition": { - "query_time_index": 1, - "conditions": [ - { - "column": "driver", - "kind": "equality", - "value": 1.5 - } - ] - }, + "conditions": [ + { + "column": "driver", + "kind": "equality", + "value": 1.5, + "query_time_indices": [1] + } + ], "seed": 7 } diff --git a/tests/fixtures/condition_mean_response.json b/tests/fixtures/condition_mean_response.json index 0c06ed6..db2d7ca 100644 --- a/tests/fixtures/condition_mean_response.json +++ b/tests/fixtures/condition_mean_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", @@ -7,9 +7,10 @@ "query_mode": "condition", "return_mode": "mean", "outputs": { - "query_times": [3], + "query_times": [2, 3], "requested_columns": ["target"], "mean": [ + [12.0], [13.5] ], "samples": null, @@ -21,7 +22,7 @@ }, "diagnostics": { "history_rows": 2, - "horizon_count": 1, + "horizon_count": 2, "seed": 7, "condition_draws": 500, "interval_estimator": null diff --git a/tests/fixtures/forecast_log_prob_request.json b/tests/fixtures/forecast_log_prob_request.json index 6ba470f..ccd069d 100644 --- a/tests/fixtures/forecast_log_prob_request.json +++ b/tests/fixtures/forecast_log_prob_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "forecast", "return_mode": "log_prob", diff --git a/tests/fixtures/forecast_log_prob_response.json b/tests/fixtures/forecast_log_prob_response.json index ab9ee00..548207b 100644 --- a/tests/fixtures/forecast_log_prob_response.json +++ b/tests/fixtures/forecast_log_prob_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/forecast_mean_request.json b/tests/fixtures/forecast_mean_request.json index 17c29ff..18d43ac 100644 --- a/tests/fixtures/forecast_mean_request.json +++ b/tests/fixtures/forecast_mean_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "forecast", "return_mode": "mean", diff --git a/tests/fixtures/forecast_mean_response.json b/tests/fixtures/forecast_mean_response.json index d289237..dab2ecd 100644 --- a/tests/fixtures/forecast_mean_response.json +++ b/tests/fixtures/forecast_mean_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/forecast_quantiles_request.json b/tests/fixtures/forecast_quantiles_request.json index 32442ba..a1551df 100644 --- a/tests/fixtures/forecast_quantiles_request.json +++ b/tests/fixtures/forecast_quantiles_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "forecast", "return_mode": "quantiles", diff --git a/tests/fixtures/forecast_quantiles_response.json b/tests/fixtures/forecast_quantiles_response.json index cd49f12..c61c0da 100644 --- a/tests/fixtures/forecast_quantiles_response.json +++ b/tests/fixtures/forecast_quantiles_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/forecast_samples_request.json b/tests/fixtures/forecast_samples_request.json index 9203269..7b06d74 100644 --- a/tests/fixtures/forecast_samples_request.json +++ b/tests/fixtures/forecast_samples_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "forecast", "return_mode": "samples", diff --git a/tests/fixtures/forecast_samples_response.json b/tests/fixtures/forecast_samples_response.json index 29cff5c..45a10b4 100644 --- a/tests/fixtures/forecast_samples_response.json +++ b/tests/fixtures/forecast_samples_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/health_metadata.json b/tests/fixtures/health_metadata.json index 804fdce..836b1ec 100644 --- a/tests/fixtures/health_metadata.json +++ b/tests/fixtures/health_metadata.json @@ -1,6 +1,6 @@ { "status": "ok", - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/input_size_exceeded_response.json b/tests/fixtures/input_size_exceeded_response.json index 2e1c7c1..ab5be10 100644 --- a/tests/fixtures/input_size_exceeded_response.json +++ b/tests/fixtures/input_size_exceeded_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v4", + "schema_version": "v5", "errors": [ { "code": "INPUT_SIZE_EXCEEDED", diff --git a/tests/fixtures/model_version_mismatch_response.json b/tests/fixtures/model_version_mismatch_response.json index 07f5363..321c7fa 100644 --- a/tests/fixtures/model_version_mismatch_response.json +++ b/tests/fixtures/model_version_mismatch_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v4", + "schema_version": "v5", "errors": [ { "code": "MODEL_VERSION_MISMATCH", diff --git a/tests/fixtures/schema_version_mismatch_response.json b/tests/fixtures/schema_version_mismatch_response.json index d5cffa9..5cd59f7 100644 --- a/tests/fixtures/schema_version_mismatch_response.json +++ b/tests/fixtures/schema_version_mismatch_response.json @@ -1,9 +1,9 @@ { - "schema_version": "v3", + "schema_version": "v4", "errors": [ { "code": "SCHEMA_VERSION_MISMATCH", - "message": "Unsupported schema_version: expected 'v4', got 'v3'", + "message": "Unsupported schema_version: expected 'v5', got 'v4'", "field": "schema_version", "detail": {} } diff --git a/tests/fixtures/validation_error_response.json b/tests/fixtures/validation_error_response.json index 30d6160..226105d 100644 --- a/tests/fixtures/validation_error_response.json +++ b/tests/fixtures/validation_error_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v4", + "schema_version": "v5", "errors": [ { "code": "VALIDATION_ERROR", diff --git a/tests/test_cli.py b/tests/test_cli.py index 006392a..1cc1d98 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -48,7 +48,7 @@ class FakeHealthClient: "deployment-id/predictionsUnstructured" ), deployment_selector="deployment_id", - schema_version="v4", + schema_version="v5", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -58,7 +58,7 @@ def health(self, *, cache: bool = False, refresh: bool = False) -> HealthMetadat del cache, refresh return HealthMetadata( status="ok", - schema_version="v4", + schema_version="v5", image_version="0.3.0", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", checkpoint_version="sdk-test", @@ -172,7 +172,7 @@ def test_predict_command_writes_response_file(monkeypatch, tmp_path: Path) -> No request_file.write_text( json.dumps( { - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", } ), diff --git a/tests/test_condition_mode.py b/tests/test_condition_mode.py index 7a1e25c..50557aa 100644 --- a/tests/test_condition_mode.py +++ b/tests/test_condition_mode.py @@ -15,8 +15,9 @@ """Tests for the ``condition`` query mode on the client side. Three surfaces are covered, and the split matters. The condition objects -validate what is wrong independently of any deployment, so a caller learns it -without a round trip. The capability gate refuses what *this* deployment has +validate what is wrong independently of any deployment — including two +conditions on one column at a shared position — so a caller learns it without a +round trip. The capability gate refuses what *this* deployment has not advertised, which is the only way to know before paying for a request. The response parser reads back the two numbers the service reports and never refuses on, so acting on them stays the caller's decision. @@ -31,7 +32,7 @@ from jointfm_client import ( ColumnSpec, - ConditionBlock, + Condition, ConditionPlausibility, DataFrameSchema, EqualityCondition, @@ -46,9 +47,10 @@ build_forecast_payload, build_forecast_payload_from_dataframe, require_condition_support, + resolve_conditions, validate_service_metadata, ) -from jointfm_client.contract import QueryMode +from jointfm_client.contract import QueryMode, condition_kinds from jointfm_client.exceptions import UnsupportedServiceContractError _MODEL_VERSION = "jointfm-inference:0.3.0+ckpt.sdk-test" @@ -67,12 +69,12 @@ def _schema() -> DataFrameSchema: def _request( - block: ConditionBlock | None, + condition: Condition | list[Condition] | None, *, requested_columns: tuple[str, ...] | None = ("target",), query_mode: QueryMode = "condition", ) -> ForecastRequest: - """Build one condition request against the two-column schema.""" + """Build one condition request against the three-column schema.""" return ForecastRequest( metadata=ForecastRequestMetadata( model_version=_MODEL_VERSION, @@ -85,7 +87,7 @@ def _request( ), query_times=(2, 3), requested_columns=requested_columns, - condition=block, + condition=condition, ) @@ -97,7 +99,7 @@ def _health( """Build one advertisement without going through a transport.""" return HealthMetadata( status="ok", - schema_version="v4", + schema_version="v5", image_version="0.3.0", model_version=_MODEL_VERSION, checkpoint_version="sdk-test", @@ -136,108 +138,184 @@ def test_an_interval_condition_rejects_a_range_that_bounds_nothing( IntervalCondition(column="driver", lower=lower, upper=upper) -def test_a_block_rejects_two_conditions_on_one_column() -> None: - """A column carries at most one condition, of either kind.""" - with pytest.raises(ValueError, match="more than one condition"): - ConditionBlock( - query_time_index=0, - conditions=( - EqualityCondition(column="driver", value=1.0), - IntervalCondition(column="driver", lower=0.0, upper=1.0), - ), +@pytest.mark.parametrize( + ("positions", "message"), + [ + ([], "must not be empty"), + ([0, 0], "more than once"), + ([-1], "must not be negative"), + ([True], "must hold integers"), + ("0", "JSON array"), + ], + ids=["empty", "duplicate", "negative", "bool", "string"], +) +def test_a_condition_rejects_positions_that_name_nothing_usable( + positions: Any, message: str +) -> None: + """``None`` covers every position; an explicit list must name real, distinct ones.""" + with pytest.raises(ValueError, match=message): + EqualityCondition(column="driver", value=1.0, query_time_indices=positions) + + +def test_a_condition_keeps_its_positions_as_an_immutable_tuple() -> None: + """A frozen condition must not change because the caller's list did.""" + positions = [0, 2] + condition = EqualityCondition( + column="driver", value=1.0, query_time_indices=positions + ) + positions.append(1) + + assert condition.query_time_indices == (0, 2) + + +def test_one_condition_and_a_list_of_them_resolve_alike() -> None: + """A single condition is the one-element list, so callers need not wrap it.""" + pin = EqualityCondition(column="driver", value=1.0) + + assert resolve_conditions(pin) == (pin,) + assert resolve_conditions([pin]) == (pin,) + + +@pytest.mark.parametrize( + ("first", "second", "overlap"), + [ + ((0,), (0, 1), r"query_time_indices \[0\]"), + (None, (1,), "every position"), + (None, None, "every position"), + ], + ids=["explicit", "null_against_explicit", "null_against_null"], +) +def test_two_conditions_on_one_column_at_a_shared_position_are_refused( + first: tuple[int, ...] | None, second: tuple[int, ...] | None, overlap: str +) -> None: + """A column carries one condition per position, whatever the two kinds are.""" + with pytest.raises( + ValueError, + match=rf"'driver' carries more than one condition at {overlap}: " + r"condition\[1\] overlaps condition\[0\]", + ): + resolve_conditions( + [ + EqualityCondition(column="driver", value=1.0, query_time_indices=first), + IntervalCondition( + column="driver", lower=0.0, upper=2.0, query_time_indices=second + ), + ] ) -def test_a_block_reports_the_kinds_a_deployment_must_advertise() -> None: - """The gate needs the kinds, and each kind appears once however many columns use it.""" - block = ConditionBlock( - query_time_index=0, - conditions=( +def test_one_column_may_carry_conditions_at_disjoint_positions() -> None: + """Disjoint positions never meet, so each keeps its own condition.""" + conditions = resolve_conditions( + [ + EqualityCondition(column="driver", value=1.0, query_time_indices=[0]), + EqualityCondition(column="driver", value=2.0, query_time_indices=[1]), + ] + ) + + assert len(conditions) == 2 + + +def test_different_columns_may_share_every_position() -> None: + """A pin and a band on different columns meet at a position; that is the point.""" + conditions = resolve_conditions( + [ EqualityCondition(column="driver", value=1.0), - IntervalCondition(column="target", lower=0.0, upper=None), - ), + IntervalCondition(column="hedge", lower=0.0, upper=None), + ] ) - assert block.kinds == ("equality", "interval") - assert block.pinned_columns == ("driver",) - assert block.conditioned_columns == ("driver", "target") + assert condition_kinds(conditions) == ("equality", "interval") -def test_the_mode_and_the_block_must_agree() -> None: +@pytest.mark.parametrize( + ("condition", "message"), + [([], "must not be empty"), (["driver"], "must be an EqualityCondition")], + ids=["empty", "not_a_condition"], +) +def test_a_condition_list_must_hold_conditions(condition: Any, message: str) -> None: + """An empty list conditions nothing, and only the two condition kinds exist.""" + with pytest.raises(ValueError, match=message): + resolve_conditions(condition) + + +def test_the_mode_and_the_condition_must_agree() -> None: """One without the other means a different request than the caller wrote.""" - block = ConditionBlock( - query_time_index=0, - conditions=(EqualityCondition(column="driver", value=1.0),), - ) + pin = EqualityCondition(column="driver", value=1.0) with pytest.raises(ValueError, match="go together"): - _request(block, query_mode="forecast") + _request(pin, query_mode="forecast") with pytest.raises(ValueError, match="go together"): _request(None) @pytest.mark.parametrize( - ("block", "requested_columns", "message"), + ("condition", "requested_columns", "message"), [ ( - ConditionBlock( - query_time_index=0, - conditions=(EqualityCondition(column="absent", value=1.0),), - ), + EqualityCondition(column="absent", value=1.0), ("target",), "undeclared columns", ), ( - ConditionBlock( - query_time_index=9, - conditions=(EqualityCondition(column="driver", value=1.0),), - ), + EqualityCondition(column="driver", value=1.0, query_time_indices=[1, 9]), ("target",), - "outside the 2 requested future positions", + r"query_time_indices \[9\] outside the 2 requested future positions", ), ( - ConditionBlock( - query_time_index=0, - conditions=( - EqualityCondition(column="driver", value=1.0), - EqualityCondition(column="hedge", value=0.5), - EqualityCondition(column="target", value=2.0), - ), - ), + [ + EqualityCondition(column="driver", value=1.0), + EqualityCondition(column="hedge", value=0.5), + EqualityCondition(column="target", value=2.0, query_time_indices=[1]), + ], None, - "log_prob", + "pinning every column at query_time_indices position 1.*log_prob", ), ( - ConditionBlock( - query_time_index=0, - conditions=( - EqualityCondition(column="driver", value=1.0), - EqualityCondition(column="hedge", value=0.5), - IntervalCondition(column="target", lower=0.0, upper=None), + [ + EqualityCondition(column="driver", value=1.0, query_time_indices=[0]), + EqualityCondition(column="hedge", value=0.5, query_time_indices=[0]), + IntervalCondition( + column="target", lower=0.0, upper=None, query_time_indices=[0] ), - ), + ], None, - "at least one column unconditioned", + "at least one column unconditioned at query_time_indices position 0", ), ], + ids=["undeclared", "outside", "all_pinned", "all_conditioned"], ) def test_a_request_is_checked_against_its_own_schema_before_any_round_trip( - block: ConditionBlock, requested_columns: tuple[str, ...] | None, message: str + condition: Condition | list[Condition], + requested_columns: tuple[str, ...] | None, + message: str, ) -> None: - """The caller learns which column name it got wrong without paying for a request.""" + """The caller learns which column or position it got wrong without a request.""" with pytest.raises(ValueError, match=message): - _request(block, requested_columns=requested_columns) + _request(condition, requested_columns=requested_columns) + + +def test_every_column_may_be_conditioned_somewhere_if_never_all_at_once() -> None: + """The read-out rule holds per position, not over the request as a whole.""" + request = _request( + [ + EqualityCondition(column="driver", value=1.0), + EqualityCondition(column="hedge", value=0.5, query_time_indices=[0]), + EqualityCondition(column="target", value=2.0, query_time_indices=[1]), + ], + requested_columns=None, + ) + + assert len(request.to_payload()["conditions"]) == 3 def test_an_interval_column_may_still_be_read_out() -> None: """An interval fixes a region, so the column's distribution inside it is an answer.""" - block = ConditionBlock( - query_time_index=0, - conditions=(IntervalCondition(column="driver", lower=0.5, upper=1.5),), + request = _request( + IntervalCondition(column="driver", lower=0.5, upper=1.5), + requested_columns=("driver", "target"), ) - request = _request(block, requested_columns=("driver", "target")) - assert request.to_payload()["requested_columns"] == ["driver", "target"] @@ -247,13 +325,11 @@ def test_a_pinned_column_may_be_read_out_in_any_position() -> None: Its answer is the request's own value, which is what lets a scenario frame be compared against an unconditioned one column for column. """ - block = ConditionBlock( - query_time_index=0, - conditions=(EqualityCondition(column="driver", value=1.0),), + request = _request( + EqualityCondition(column="driver", value=1.0), + requested_columns=("target", "driver"), ) - request = _request(block, requested_columns=("target", "driver")) - assert request.to_payload()["requested_columns"] == ["target", "driver"] @@ -263,18 +339,15 @@ def test_omitting_the_projection_states_nothing_on_the_wire() -> None: A client-side default would be a second copy of the rule, free to drift from the one the deployment actually applies. """ - block = ConditionBlock( - query_time_index=0, - conditions=(EqualityCondition(column="driver", value=1.0),), - ) - - payload = _request(block, requested_columns=None).to_payload() + payload = _request( + EqualityCondition(column="driver", value=1.0), requested_columns=None + ).to_payload() assert "requested_columns" not in payload -def test_the_payload_carries_the_block_the_service_parses() -> None: - """The wire form names its position once and each condition names its kind.""" +def test_the_payload_carries_the_conditions_the_service_parses() -> None: + """Each condition names its kind and its positions, ``null`` covering them all.""" payload = build_forecast_payload( model_version=_MODEL_VERSION, schema=_schema(), @@ -282,51 +355,56 @@ def test_the_payload_carries_the_block_the_service_parses() -> None: query_times=(2, 3), requested_columns=("target",), query_mode="condition", - condition=ConditionBlock( - query_time_index=1, - conditions=( - EqualityCondition(column="driver", value=1.5), - IntervalCondition(column="target", lower=None, upper=12.0), + condition=[ + EqualityCondition(column="driver", value=1.5), + IntervalCondition( + column="target", lower=None, upper=12.0, query_time_indices=[1] ), - ), + ], ) assert payload["query_mode"] == "condition" - assert payload["condition"] == { - "query_time_index": 1, - "conditions": [ - {"column": "driver", "kind": "equality", "value": 1.5}, - {"column": "target", "kind": "interval", "lower": None, "upper": 12.0}, - ], - } + assert "condition" not in payload + assert payload["conditions"] == [ + { + "column": "driver", + "kind": "equality", + "value": 1.5, + "query_time_indices": None, + }, + { + "column": "target", + "kind": "interval", + "lower": None, + "upper": 12.0, + "query_time_indices": [1], + }, + ] def test_the_gate_refuses_a_deployment_that_advertises_no_condition_mode() -> None: """Discovery before the request is the point: the error must not cost a round trip.""" - block = ConditionBlock( - query_time_index=0, - conditions=(EqualityCondition(column="driver", value=1.0),), - ) + pin = EqualityCondition(column="driver", value=1.0) with pytest.raises( UnsupportedServiceContractError, match="does not serve condition" ): require_condition_support( - _health(query_modes=("forecast",), condition_kinds=()), block + _health(query_modes=("forecast",), condition_kinds=()), pin ) def test_the_gate_refuses_a_kind_the_deployment_does_not_answer() -> None: - """A deployment may serve one kind before the other; the block says which it needs.""" - block = ConditionBlock( - query_time_index=0, - conditions=(IntervalCondition(column="driver", lower=0.0, upper=1.0),), - ) + """A deployment may serve one kind before the other; the conditions say which.""" + conditions = [ + EqualityCondition(column="driver", value=1.0), + IntervalCondition(column="hedge", lower=0.0, upper=1.0), + ] with pytest.raises(UnsupportedServiceContractError, match=r"kinds \['interval'\]"): - require_condition_support(_health(condition_kinds=("equality",)), block) + require_condition_support(_health(condition_kinds=("equality",)), conditions) - require_condition_support(_health(), block) + require_condition_support(_health(), conditions) def test_the_equality_response_reads_back_its_plausibility( @@ -348,20 +426,37 @@ def test_the_equality_response_reads_back_its_plausibility( assert response.diagnostics.interval_estimator is None -def test_the_response_describes_the_conditioned_position_alone( +def test_the_response_answers_every_requested_position( json_fixture_loader: Callable[[str], dict[str, Any]], ) -> None: - """The request asked for two future rows; a condition answers about one.""" + """The condition covers one of two rows, and the answer still carries both.""" request_payload = json_fixture_loader("condition_mean_request") assert request_payload["query_times"] == [2, 3] + assert request_payload["conditions"][0]["query_time_indices"] == [1] response = ForecastResponse.from_payload( json_fixture_loader("condition_mean_response"), request_payload=request_payload, ) - assert response.query_times == (3,) - assert response.diagnostics.horizon_count == 1 + assert response.query_times == (2, 3) + assert response.diagnostics.horizon_count == 2 + + +def test_a_response_narrowed_to_the_covered_position_is_refused( + json_fixture_loader: Callable[[str], dict[str, Any]], +) -> None: + """A deployment answering fewer positions than asked must fail, not parse.""" + response_payload = json_fixture_loader("condition_mean_response") + response_payload["outputs"]["query_times"] = [3] + response_payload["outputs"]["mean"] = [[13.5]] + response_payload["diagnostics"]["horizon_count"] = 1 + + with pytest.raises(ValueError, match="query_times"): + ForecastResponse.from_payload( + response_payload, + request_payload=json_fixture_loader("condition_mean_request"), + ) def test_an_interval_response_reads_back_its_estimator_accuracy( @@ -418,7 +513,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: assert isinstance(sample_count, int) start = sum(cast(int, earlier["n_samples"]) for earlier in self.payloads[:-1]) return { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": _MODEL_VERSION, "checkpoint_version": "sdk-test", @@ -426,10 +521,12 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: "query_mode": "condition", "return_mode": "samples", "outputs": { - "query_times": [3], + "query_times": [2, 3], "requested_columns": ["target"], "mean": None, - "samples": [[[float(start + index)]] for index in range(sample_count)], + "samples": [ + [[-1.0], [float(start + index)]] for index in range(sample_count) + ], "quantiles": None, }, "plausibility": { @@ -438,7 +535,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: }, "diagnostics": { "history_rows": 2, - "horizon_count": 1, + "horizon_count": 2, "seed": payload.get("seed"), "condition_draws": sample_count, "interval_estimator": None, @@ -456,7 +553,7 @@ def _health_payload( """Build one health advertisement as the service serializes it.""" return { "status": "ok", - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": _MODEL_VERSION, "checkpoint_version": "sdk-test", @@ -490,16 +587,13 @@ def _client(transport: _ConditionTransport) -> JointFMClient: {"driver": 1.0, "hedge": 0.5, "target": 10.0}, {"driver": 1.1, "hedge": 0.6, "target": 11.0}, ) -_PIN_DRIVER = ConditionBlock( - query_time_index=1, - conditions=(EqualityCondition(column="driver", value=1.5),), -) +_PIN_DRIVER = EqualityCondition(column="driver", value=1.5, query_time_indices=[1]) def test_the_client_sends_the_condition_and_reads_the_answer_back( json_fixture_loader: Callable[[str], dict[str, Any]], ) -> None: - """The typed helper carries the block onto the wire and types what comes back.""" + """The typed helper carries the condition onto the wire and types what comes back.""" transport = _ConditionTransport( health_payload=_health_payload( query_modes=["forecast", "condition"], @@ -520,13 +614,13 @@ def test_the_client_sends_the_condition_and_reads_the_answer_back( assert isinstance(result, MeanForecastResult) assert result.query_mode == "condition" - assert result.query_times == (3,) - assert result.mean == ((13.5,),) + assert result.query_times == (2, 3) + assert result.mean == ((12.0,), (13.5,)) assert result.plausibility == ConditionPlausibility(equality_log_density=-1.27) assert len(transport.payloads) == 1 sent = transport.payloads[0] assert sent["query_mode"] == "condition" - assert sent["condition"] == _PIN_DRIVER.to_payload() + assert sent["conditions"] == [_PIN_DRIVER.to_payload()] def test_the_client_refuses_before_posting_when_the_deployment_cannot_condition() -> ( @@ -552,7 +646,7 @@ def test_the_client_refuses_before_posting_when_the_deployment_cannot_condition( def test_batched_condition_samples_merge_into_one_conditional_answer() -> None: - """Every batch repeats the same block; the merge recounts the draws and keeps the plausibility.""" + """Every batch repeats the same conditions; the merge recounts draws, keeps plausibility.""" transport = _ConditionTransport( health_payload=_health_payload( query_modes=["forecast", "condition"], @@ -573,20 +667,24 @@ def test_batched_condition_samples_merge_into_one_conditional_answer() -> None: ) assert isinstance(result, SampleForecastResult) - assert result.samples == (((0.0,),), ((1.0,),), ((2.0,),)) - assert result.query_times == (3,) + assert result.samples == ( + ((-1.0,), (0.0,)), + ((-1.0,), (1.0,)), + ((-1.0,), (2.0,)), + ) + assert result.query_times == (2, 3) assert result.diagnostics.condition_draws == 3 assert result.plausibility == ConditionPlausibility(equality_log_density=-1.27) assert [payload["n_samples"] for payload in transport.payloads] == [2, 1] assert all( payload["query_mode"] == "condition" - and payload["condition"] == _PIN_DRIVER.to_payload() + and payload["conditions"] == [_PIN_DRIVER.to_payload()] for payload in transport.payloads ) def test_the_dataframe_adapter_builds_a_condition_request() -> None: - """A pandas caller passes the block and gets the condition envelope.""" + """A pandas caller passes the conditions and gets the condition envelope.""" pandas = pytest.importorskip("pandas") frame = pandas.DataFrame(list(_HISTORY_ROWS)) @@ -601,7 +699,7 @@ def test_the_dataframe_adapter_builds_a_condition_request() -> None: ) assert payload["query_mode"] == "condition" - assert payload["condition"] == _PIN_DRIVER.to_payload() + assert payload["conditions"] == [_PIN_DRIVER.to_payload()] assert payload["requested_columns"] == ["target"] diff --git a/tests/test_configuration.py b/tests/test_configuration.py index 3a6fe7b..88344fe 100644 --- a/tests/test_configuration.py +++ b/tests/test_configuration.py @@ -47,7 +47,7 @@ class _HealthTransport: _METADATA: dict[str, object] = { "status": "ok", - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.yaml", "checkpoint_version": "yaml", @@ -147,7 +147,7 @@ def test_load_settings_layers_config_below_dotenv_and_environment( "datarobot_endpoint": "https://app.datarobot.com/api/v2", "datarobot_api_token": "yaml-token", "deployment_id": "yaml-deployment-id", - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.yaml", } }, @@ -191,7 +191,7 @@ def test_client_from_env_uses_transport_defaults_from_config( "datarobot_endpoint": "https://app.datarobot.com/api/v2", "datarobot_api_token": "yaml-token", "deployment_id": "yaml-deployment-id", - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.yaml", }, "transport": { diff --git a/tests/test_contract.py b/tests/test_contract.py index 9d80ad3..30ee0a1 100644 --- a/tests/test_contract.py +++ b/tests/test_contract.py @@ -76,7 +76,7 @@ def _health_metadata() -> dict[str, object]: """Health metadata.""" return { "status": "ok", - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -111,7 +111,7 @@ def test_package_identity_contract() -> None: assert DISTRIBUTION_NAME == "jointfm-client" assert IMPORT_NAMESPACE == "jointfm_client" assert FIRST_SUPPORTED_PYTHON_VERSION == "3.11" - assert SCHEMA_VERSION == "v4" + assert SCHEMA_VERSION == "v5" def test_package_version_matches_installed_distribution() -> None: @@ -267,7 +267,7 @@ def test_forecast_payload_matches_service_contract_without_mutating_inputs() -> ) assert payload == { - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "quantiles", @@ -362,7 +362,7 @@ def test_dataframe_payload_matches_service_forecast_request_shape() -> None: ) assert payload == { - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "mean", @@ -777,7 +777,7 @@ def test_health_and_response_models_parse_current_payloads() -> None: health = HealthMetadata.from_payload(_health_metadata()) response = ForecastResponse.from_payload( { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -840,7 +840,7 @@ def test_forecast_result_conversion_helpers_cover_mean_samples_and_quantiles() - """Forecast result conversion helpers cover mean samples and quantiles.""" mean_result = ForecastResponse.from_payload( { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -860,7 +860,7 @@ def test_forecast_result_conversion_helpers_cover_mean_samples_and_quantiles() - ) sample_result = ForecastResponse.from_payload( { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -883,7 +883,7 @@ def test_forecast_result_conversion_helpers_cover_mean_samples_and_quantiles() - ) quantile_result = ForecastResponse.from_payload( { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -957,7 +957,7 @@ def test_forecast_result_conversion_helpers_cover_mean_samples_and_quantiles() - def test_forecast_response_validates_request_scoped_shapes() -> None: """Forecast response validates request scoped shapes.""" request_payload = { - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "samples", @@ -966,7 +966,7 @@ def test_forecast_response_validates_request_scoped_shapes() -> None: "n_samples": 3, } response_payload = { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -996,7 +996,7 @@ def test_forecast_response_raises_typed_error_for_success_payload_errors() -> No with pytest.raises(JointFMServiceError) as exc_info: ForecastResponse.from_payload( { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -1162,7 +1162,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: def _forecast_response_payload(*, return_mode: str) -> dict[str, object]: """Forecast response payload.""" return { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", diff --git a/tests/test_contract_models.py b/tests/test_contract_models.py index af48ec8..7352acd 100644 --- a/tests/test_contract_models.py +++ b/tests/test_contract_models.py @@ -76,7 +76,7 @@ def test_request_models_serialize_direct_payloads_without_mutating_inputs() -> N ) assert metadata.to_payload() == { - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "quantiles", @@ -91,7 +91,7 @@ def test_request_models_serialize_direct_payloads_without_mutating_inputs() -> N "timezone": "UTC", } assert request.to_payload() == { - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "quantiles", @@ -303,7 +303,7 @@ def test_response_models_reject_direct_validation_edges() -> None: def test_forecast_response_rejects_request_scoped_metadata_mismatches() -> None: """Forecast response rejects request scoped metadata mismatches.""" request_payload = { - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "mean", @@ -354,7 +354,7 @@ def test_forecast_response_rejects_request_scoped_metadata_mismatches() -> None: def test_forecast_response_rejects_sample_bound_violations() -> None: """Forecast response rejects sample bound violations.""" request_payload = { - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "samples", @@ -397,7 +397,7 @@ def test_forecast_response_rejects_sample_bound_violations() -> None: def _mean_response_payload() -> dict[str, Any]: """Mean response payload.""" return { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -419,7 +419,7 @@ def _mean_response_payload() -> dict[str, Any]: def _sample_response_payload() -> dict[str, Any]: """Sample response payload.""" return { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", diff --git a/tests/test_feature_importance.py b/tests/test_feature_importance.py index c706175..273d84a 100644 --- a/tests/test_feature_importance.py +++ b/tests/test_feature_importance.py @@ -47,7 +47,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: self.payloads.append(dict(payload)) samples = self.sample_batches[len(self.payloads) - 1] return { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": _MODEL_VERSION, "checkpoint_version": "sdk-test", diff --git a/tests/test_log_prob_mode.py b/tests/test_log_prob_mode.py index 9c0cd34..2715b53 100644 --- a/tests/test_log_prob_mode.py +++ b/tests/test_log_prob_mode.py @@ -35,7 +35,7 @@ from jointfm_client import ( ColumnSpec, - ConditionBlock, + Condition, ConditionPlausibility, DataFrameSchema, EqualityCondition, @@ -61,10 +61,7 @@ {"driver": 1.2, "target": 12.0}, {"driver": 1.5, "target": 13.0}, ) -_PIN_DRIVER = ConditionBlock( - query_time_index=1, - conditions=(EqualityCondition(column="driver", value=1.5),), -) +_PIN_DRIVER = EqualityCondition(column="driver", value=1.5, query_time_indices=[1]) def _schema() -> DataFrameSchema: @@ -83,7 +80,7 @@ def _request( return_mode: ReturnMode = "log_prob", query_rows: Any = _QUERY_ROWS, requested_columns: tuple[str, ...] | None = ("driver", "target"), - condition: ConditionBlock | None = None, + condition: Condition | None = None, ) -> ForecastRequest: """Build one scoring request against the two-column schema.""" return ForecastRequest( @@ -218,7 +215,7 @@ def test_a_score_may_leave_out_no_column_at_all_under_a_condition() -> None: def test_the_payload_carries_the_rows_the_service_scores( json_fixture_loader: Callable[[str], dict[str, Any]], fixture_name: str, - condition: ConditionBlock | None, + condition: Condition | None, requested_columns: tuple[str, ...], query_rows: tuple[Mapping[str, Any], ...], ) -> None: @@ -324,7 +321,7 @@ def test_the_client_scores_observed_rows_and_types_the_answer( def test_the_client_scores_under_a_condition( json_fixture_loader: Callable[[str], dict[str, Any]], ) -> None: - """A conditioned score answers at the conditioned position and reports its pin.""" + """A conditioned score answers every position and reports its pin.""" health_payload = json_fixture_loader("health_metadata") health_payload["supported_query_modes"] = ["forecast", "condition"] health_payload["supported_condition_kinds"] = ["equality", "interval"] @@ -345,8 +342,8 @@ def test_the_client_scores_under_a_condition( ) assert isinstance(result, LogProbResult) - assert result.query_times == (3,) - assert result.log_prob.values == (-1.75,) + assert result.query_times == (2, 3) + assert result.log_prob.values == (-2.5, -1.75) assert result.plausibility == ConditionPlausibility(equality_log_density=-1.27) sent = transport.payloads[0] assert sent["query_mode"] == "condition" diff --git a/tests/test_pool.py b/tests/test_pool.py index 9de17bb..7615efd 100644 --- a/tests/test_pool.py +++ b/tests/test_pool.py @@ -56,7 +56,7 @@ def _health( """Health.""" return { "status": "ok", - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": model_version, "checkpoint_version": checkpoint_version, @@ -134,7 +134,7 @@ def test_pool_retries_next_instance_on_470() -> None: pool = _pool(transport=_Transport(fail_ids=frozenset({"a"}))) assert pool.next_instance().deployment_id == "a" assert pool.next_instance().deployment_id == "b" - assert pool.post_json({"schema_version": "v4"}) == { + assert pool.post_json({"schema_version": "v5"}) == { "ok": True, "deployment_id": "b", } @@ -144,7 +144,7 @@ def test_pool_raises_when_all_instances_unavailable() -> None: """Pool raises when all instances unavailable.""" pool = _pool(transport=_Transport(fail_ids=frozenset({"a", "b"}))) with pytest.raises(JointFMHTTPStatusError, match="unavailable"): - pool.post_json({"schema_version": "v4"}) + pool.post_json({"schema_version": "v5"}) def test_pool_health_rejects_mismatch_and_aligns_sample_cap() -> None: @@ -275,7 +275,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: executor.submit( pool.post_json_to, pool.instance_at(index), - {"schema_version": "v4"}, + {"schema_version": "v5"}, ) for index in range(2) ] @@ -299,7 +299,7 @@ def test_pool_health_routes_only_reachable_peers() -> None: assert pool.instance_at(0).deployment_id == "b" assert pool.instance_at(1).deployment_id == "b" assert pool.next_instance().deployment_id == "b" - assert pool.post_json({"schema_version": "v4"})["deployment_id"] == "b" + assert pool.post_json({"schema_version": "v5"})["deployment_id"] == "b" def test_pool_failover_retries_health_excluded_peer() -> None: @@ -313,7 +313,7 @@ def test_pool_failover_retries_health_excluded_peer() -> None: assert pool.instance_at(0).deployment_id == "b" transport.fail_ids = frozenset({"b"}) - assert pool.post_json({"schema_version": "v4"})["deployment_id"] == "a" + assert pool.post_json({"schema_version": "v5"})["deployment_id"] == "a" assert pool.instance_at(0).deployment_id == "a" @@ -333,7 +333,7 @@ def test_pool_health_skips_incompatible_peer_when_another_matches_pin() -> None: assert metadata.model_version == pinned assert pool.instance_at(0).deployment_id == "a" assert pool.instance_at(1).deployment_id == "a" - assert pool.post_json({"schema_version": "v4"})["deployment_id"] == "a" + assert pool.post_json({"schema_version": "v5"})["deployment_id"] == "a" def test_pool_cooldown_restores_peer_after_transient_failure( @@ -345,7 +345,7 @@ def test_pool_cooldown_restores_peer_after_transient_failure( transport = _Transport(fail_ids=frozenset({"a"})) pool = _pool(transport=transport, peer_cooldown_seconds=10.0) - assert pool.post_json({"schema_version": "v4"})["deployment_id"] == "b" + assert pool.post_json({"schema_version": "v5"})["deployment_id"] == "b" assert pool.instance_at(0).deployment_id == "b" assert pool.instance_at(1).deployment_id == "b" @@ -355,4 +355,4 @@ def test_pool_cooldown_restores_peer_after_transient_failure( clock["now"] = 110.0 assert {pool.instance_at(i).deployment_id for i in range(2)} == {"a", "b"} - assert pool.post_json({"schema_version": "v4"})["deployment_id"] in {"a", "b"} + assert pool.post_json({"schema_version": "v5"})["deployment_id"] in {"a", "b"} diff --git a/tests/test_settings.py b/tests/test_settings.py index 13632c0..2769036 100644 --- a/tests/test_settings.py +++ b/tests/test_settings.py @@ -54,7 +54,7 @@ def _hosted_env(**overrides: str) -> dict[str, str]: DATAROBOT_ENDPOINT_ENV: "https://app.datarobot.com/api/v2/", DATAROBOT_API_TOKEN_ENV: "secret-token", JOINTFM_DEPLOYMENT_ID_ENV: "deployment-id", - JOINTFM_SCHEMA_VERSION_ENV: "v4", + JOINTFM_SCHEMA_VERSION_ENV: "v5", JOINTFM_MODEL_VERSION_ENV: "jointfm-inference:0.3.0+ckpt.sdk-test", } env.update(overrides) @@ -66,7 +66,7 @@ def test_load_settings_from_environment_with_deployment_id_builds_hosted_url() - settings = load_settings(env=_hosted_env(), dotenv_path=None) assert settings.datarobot_endpoint == "https://app.datarobot.com/api/v2" - assert settings.schema_version == "v4" + assert settings.schema_version == "v5" assert settings.model_version == "jointfm-inference:0.3.0+ckpt.sdk-test" assert settings.deployment_id == "deployment-id" assert settings.predict_url == ( @@ -82,7 +82,7 @@ def test_load_settings_with_local_service_base_url_builds_direct_urls() -> None: settings = load_settings( env={ JOINTFM_LOCAL_BASE_URL_ENV: "http://127.0.0.1:8080/", - JOINTFM_SCHEMA_VERSION_ENV: "v4", + JOINTFM_SCHEMA_VERSION_ENV: "v5", JOINTFM_MODEL_VERSION_ENV: "jointfm-inference:0.3.0+ckpt.local-test", }, dotenv_path=None, @@ -158,7 +158,7 @@ def test_load_settings_reads_dotenv_without_overriding_environment(tmp_path) -> "DATAROBOT_ENDPOINT=https://app.datarobot.com/api/v2", "DATAROBOT_API_TOKEN=file-token", "JOINTFM_DEPLOYMENT_ID=file-deployment-id", - "JOINTFM_SCHEMA_VERSION=v4", + "JOINTFM_SCHEMA_VERSION=v5", "JOINTFM_MODEL_VERSION=jointfm-inference:0.3.0+ckpt.sdk-test", ] ), @@ -209,7 +209,7 @@ def test_load_settings_rejects_missing_credentials_without_defaults() -> None: env={ DATAROBOT_API_TOKEN_ENV: "secret-token", JOINTFM_DEPLOYMENT_ID_ENV: "deployment-id", - JOINTFM_SCHEMA_VERSION_ENV: "v4", + JOINTFM_SCHEMA_VERSION_ENV: "v5", JOINTFM_MODEL_VERSION_ENV: "jointfm-inference:0.3.0+ckpt.sdk-test", }, dotenv_path=None, @@ -220,7 +220,7 @@ def test_load_settings_rejects_missing_credentials_without_defaults() -> None: env={ DATAROBOT_ENDPOINT_ENV: "https://app.datarobot.com/api/v2", JOINTFM_DEPLOYMENT_ID_ENV: "deployment-id", - JOINTFM_SCHEMA_VERSION_ENV: "v4", + JOINTFM_SCHEMA_VERSION_ENV: "v5", JOINTFM_MODEL_VERSION_ENV: "jointfm-inference:0.3.0+ckpt.sdk-test", }, dotenv_path=None, diff --git a/tests/test_surfaces.py b/tests/test_surfaces.py index bfd2a26..4f79ca6 100644 --- a/tests/test_surfaces.py +++ b/tests/test_surfaces.py @@ -70,7 +70,7 @@ def test_hosted_surface_uses_datarobot_routes_and_auth_headers( health_url=predict_url, predict_url=predict_url, deployment_selector="deployment_id", - schema_version="v4", + schema_version="v5", model_version=request_payload["model_version"], deployment_id="deployment-id", ) @@ -209,7 +209,7 @@ def test_hosted_surface_auto_discovers_model_version_when_settings_unpinned( health_url=predict_url, predict_url=predict_url, deployment_selector="deployment_id", - schema_version="v4", + schema_version="v5", deployment_id="deployment-id", ) assert settings.model_version is None diff --git a/tests/test_transport.py b/tests/test_transport.py index c5fd1af..d291f99 100644 --- a/tests/test_transport.py +++ b/tests/test_transport.py @@ -154,7 +154,7 @@ def _health_payload( """Health payload.""" return { "status": "ok", - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": model_version, "checkpoint_version": checkpoint_version, @@ -190,7 +190,7 @@ def _forecast_response_payload(*, return_mode: str = "mean") -> dict[str, object else None, } return { - "schema_version": "v4", + "schema_version": "v5", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", @@ -234,7 +234,7 @@ def test_transport_posts_json_with_headers_timeout_and_user_agent() -> None: session.mount("https://", adapter) result = transport.post_json( - "https://example.com/predict", {"schema_version": "v4"} + "https://example.com/predict", {"schema_version": "v5"} ) assert result == {"ok": True} @@ -248,7 +248,7 @@ def test_transport_posts_json_with_headers_timeout_and_user_agent() -> None: assert adapter.kwargs[0]["timeout"] == (1.5, 2.5) request_body = request.body assert isinstance(request_body, bytes) - assert json.loads(request_body.decode("utf-8")) == {"schema_version": "v4"} + assert json.loads(request_body.decode("utf-8")) == {"schema_version": "v5"} def test_transport_from_settings_attaches_hosted_auth_headers_and_closes_session() -> ( @@ -265,7 +265,7 @@ def test_transport_from_settings_attaches_hosted_auth_headers_and_closes_session "deployment-id/predictionsUnstructured" ), deployment_selector="deployment_id", - schema_version="v4", + schema_version="v5", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -293,7 +293,7 @@ def test_transport_from_local_settings_omits_hosted_auth_headers() -> None: health_url="http://127.0.0.1:8080/healthz", predict_url="http://127.0.0.1:8080/predict", deployment_selector="local_service", - schema_version="v4", + schema_version="v5", model_version="jointfm-inference:0.3.0+ckpt.local-test", local_base_url="http://127.0.0.1:8080", ) @@ -319,7 +319,7 @@ def test_transport_retries_retryable_server_responses() -> None: retry_config=JointFMRetryConfig(max_attempts=2) ) - result = transport.post_json(_server_url(server), {"schema_version": "v4"}) + result = transport.post_json(_server_url(server), {"schema_version": "v5"}) assert result == {"ok": True} assert handler.request_count == 2 @@ -351,7 +351,7 @@ def test_transport_retries_html_bodied_gateway_errors() -> None: ), ) - result = transport.post_json(_server_url(server), {"schema_version": "v4"}) + result = transport.post_json(_server_url(server), {"schema_version": "v5"}) assert result == {"ok": True} assert handler.request_count == 2 @@ -380,7 +380,7 @@ def test_transport_raises_status_error_when_gateway_html_persists() -> None: ) with pytest.raises(JointFMHTTPStatusError) as exc_info: - transport.post_json(_server_url(server), {"schema_version": "v4"}) + transport.post_json(_server_url(server), {"schema_version": "v5"}) assert exc_info.value.status_code == HTTPStatus.BAD_GATEWAY assert "502 Bad Gateway" in exc_info.value.response_body_excerpt @@ -444,7 +444,7 @@ def request(self, *args: Any, **kwargs: Any) -> requests.Response: ) result = transport.post_json( - "https://example.com/predict", {"schema_version": "v4"} + "https://example.com/predict", {"schema_version": "v5"} ) assert result == {"ok": True} @@ -463,7 +463,7 @@ def test_status_error_carries_parsed_retry_after_seconds() -> None: ) with pytest.raises(JointFMHTTPStatusError) as exc_info: - transport.post_json(_server_url(server), {"schema_version": "v4"}) + transport.post_json(_server_url(server), {"schema_version": "v5"}) assert exc_info.value.retry_after_seconds == 0.5 assert handler.request_count == 1 @@ -481,7 +481,7 @@ def test_transport_does_not_retry_validation_errors() -> None: ) with pytest.raises(JointFMHTTPStatusError) as exc_info: - transport.post_json(_server_url(server), {"schema_version": "v4"}) + transport.post_json(_server_url(server), {"schema_version": "v5"}) assert exc_info.value.status_code == HTTPStatus.BAD_REQUEST assert exc_info.value.datarobot_request_id == "request-id-1" @@ -507,7 +507,7 @@ def test_transport_rejects_non_json_serializable_payloads() -> None: with pytest.raises(JointFMRequestEncodingError, match="JSON-serializable"): transport.post_json( "https://example.com/predict", - {"schema_version": "v4", "bad": object()}, + {"schema_version": "v5", "bad": object()}, ) @@ -572,14 +572,14 @@ def test_client_predict_uses_configured_transport_and_settings() -> None: health_url="https://app.datarobot.com/api/v2/deployments/deployment-id/healthz", predict_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", deployment_selector="deployment_id", - schema_version="v4", + schema_version="v5", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) transport = RecordingTransport() client = JointFMClient(settings=settings, transport=transport) payload = { - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", } @@ -598,7 +598,7 @@ def test_client_health_returns_typed_metadata_and_caches_only_when_requested() - health_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", predict_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", deployment_selector="deployment_id", - schema_version="v4", + schema_version="v5", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -629,7 +629,7 @@ def test_client_health_instances_returns_one_entry_for_single_endpoint() -> None health_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", predict_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", deployment_selector="deployment_id", - schema_version="v4", + schema_version="v5", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -691,7 +691,7 @@ def capture_transport( "DATAROBOT_ENDPOINT": "https://app.datarobot.com/api/v2", "DATAROBOT_API_TOKEN": "secret-token", "JOINTFM_DEPLOYMENT_ID": "deployment-id", - "JOINTFM_SCHEMA_VERSION": "v4", + "JOINTFM_SCHEMA_VERSION": "v5", "JOINTFM_MODEL_VERSION": "jointfm-inference:0.3.0+ckpt.sdk-test", }, dotenv_path=None, @@ -713,7 +713,7 @@ def test_client_health_rejects_cached_model_mismatch() -> None: health_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", predict_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", deployment_selector="deployment_id", - schema_version="v4", + schema_version="v5", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -739,7 +739,7 @@ def test_client_hosted_health_posts_request_type_health_to_predict_url() -> None health_url=predict_url, predict_url=predict_url, deployment_selector="deployment_id", - schema_version="v4", + schema_version="v5", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -763,7 +763,7 @@ def test_client_local_health_keeps_get_request_to_healthz_route() -> None: health_url="http://127.0.0.1:8080/healthz", predict_url="http://127.0.0.1:8080/predict", deployment_selector="local_service", - schema_version="v4", + schema_version="v5", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", local_base_url="http://127.0.0.1:8080", ) @@ -805,7 +805,7 @@ def test_client_forecast_builds_payload_from_rows_and_returns_typed_response() - assert result.outputs.mean == ((12.0,),) assert transport.predict_url == "http://localhost:8080/predict" assert transport.payload == { - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "forecast", "return_mode": "mean", @@ -980,7 +980,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: health_url="http://127.0.0.1:8080/healthz", predict_url="http://127.0.0.1:8080/predict", deployment_selector="local_service", - schema_version="v4", + schema_version="v5", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", local_base_url="http://127.0.0.1:8080", ) @@ -1032,13 +1032,13 @@ def test_client_predict_raises_typed_service_error_for_success_payload_errors() health_url="https://app.datarobot.com/api/v2/deployments/deployment-id/healthz", predict_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", deployment_selector="deployment_id", - schema_version="v4", + schema_version="v5", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) transport = RecordingTransport() transport.predict_payload = { - "schema_version": "v4", + "schema_version": "v5", "errors": [ { "code": "VALIDATION_ERROR", @@ -1052,7 +1052,7 @@ def test_client_predict_raises_typed_service_error_for_success_payload_errors() with pytest.raises(JointFMServiceError) as exc_info: client.predict( { - "schema_version": "v4", + "schema_version": "v5", "model_version": settings.model_version, } ) @@ -1141,7 +1141,7 @@ def do_POST(self) -> None: payload = {"ok": True} else: payload = { - "schema_version": "v4", + "schema_version": "v5", "errors": [ { "code": "VALIDATION_ERROR", @@ -1225,7 +1225,7 @@ def _pool_settings(primary: str, backup: str) -> JointFMSettings: health_url=primary, predict_url=primary, deployment_selector="deployment_ids", - schema_version="v4", + schema_version="v5", instances=( JointFMInstanceSettings(deployment_id="primary-id", predict_url=primary), JointFMInstanceSettings(deployment_id="backup-id", predict_url=backup), @@ -1267,7 +1267,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: transport = PoolTransport() client = JointFMClient(settings=settings, transport=transport) payload = { - "schema_version": "v4", + "schema_version": "v5", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", } client.predict(payload) From b7788c6701ff234f716d20ba5c62f87584949674 Mon Sep 17 00:00:00 2001 From: Stefan Hackmann Date: Thu, 24 Sep 2026 10:42:59 +0000 Subject: [PATCH 7/8] docs: Make the served conditioning documentation describe the served behavior --- README.md | 2 +- docs/api-reference.md | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/README.md b/README.md index 2b1e7f0..0182402 100644 --- a/README.md +++ b/README.md @@ -211,7 +211,7 @@ print(result.plausibility) Whether a deployment can condition depends on the mounted checkpoint's head. `/healthz` advertises `condition` in `supported_query_modes` and the kinds it answers in `supported_condition_kinds` (empty when the mode is absent). The client checks that advertisement before sending, so a deployment that cannot condition is refused with `UnsupportedServiceContractError` rather than after a paid round trip; `require_condition_support(metadata, block)` exposes the same check. -A condition response carries a `plausibility` block: `equality_log_density` is the log density the model assigns to the pinned values and `region_log_probability` the log probability it gives the interval region, each `None` when the request carried no condition of that kind. They separate a confident answer from one conditioned on something the model finds implausible; the service reports them and never refuses on them. `diagnostics.condition_draws` counts the draws behind sampled outputs, and `diagnostics.interval_estimator` reports the numerical accounting (`points`, `effective_sample_size`) when more than one column carries an interval and the region probability had to be estimated. +A condition response carries a `plausibility` block: `equality_log_density` is the log density the model assigns to the pinned values and `region_log_probability` the log probability it gives the interval region *given the pinned values*, so the two add to the plausibility of the whole condition set, each `None` when the request carried no condition of that kind. They separate a confident answer from one conditioned on something the model finds implausible; the service reports them and never refuses on them. `diagnostics.condition_draws` counts the draws behind sampled outputs, and `diagnostics.interval_estimator` reports the numerical accounting (`points`, `effective_sample_size`) when more than one column carries an interval and the region probability had to be estimated. Column descriptors support the server fields `name`, `modality`, `role`, `nullable`, `vocabulary_size`, `level_count`, `mapping`, `lower_bound`, `upper_bound`, `time_value_kind`, `time_value_scale_seconds`, `time_value_use_local_normalized_time`, `time_value_calendar_id`, and `time_value_timezone`. diff --git a/docs/api-reference.md b/docs/api-reference.md index efc7f82..b18c290 100644 --- a/docs/api-reference.md +++ b/docs/api-reference.md @@ -21,7 +21,7 @@ This reference covers the supported public Python surface exported by `jointfm_c | `EqualityCondition` | One column of the conditioned position pinned to a finite `value`. The column may be named in `requested_columns` like any other, and reads back the pinned value the request supplied. | | `IntervalCondition` | One column of the conditioned position confined to `[lower, upper]`; `None` leaves a side open, at least one side must be bounded, and `lower < upper`. The column stays readable and the response describes its distribution inside the range. | | `ConditionBlock` | Every condition of one request: `query_time_index` into the request's `query_times` and a sequence of `EqualityCondition` or `IntervalCondition` values, at most one per column. `kinds`, `pinned_columns`, and `conditioned_columns` expose what the capability gate and the request validation need. | -| `ConditionPlausibility` | What the model thinks of the conditions it was given: `equality_log_density` of the pinned values and `region_log_probability` of the interval region, each `None` when the request carried no condition of that kind. Reported by the service and never refused on. | +| `ConditionPlausibility` | What the model thinks of the conditions it was given: `equality_log_density` of the pinned values and `region_log_probability` of the interval region given the pinned values, each `None` when the request carried no condition of that kind. Reported by the service and never refused on. | | `IntervalEstimator` | Numerical accounting (`points`, `effective_sample_size`) behind a region probability that had to be estimated, which happens when more than one column carries an interval condition. | | `HealthMetadata` | Typed service-health payload with service status, schema and model versions, checkpoint metadata, device, head, `decoding_strategy`, advertised query modes, `supported_condition_kinds` (empty when the deployment cannot condition), return modes, time-index modes, time-index encoding, `max_sample_count`, and an optional `data_generation` block carrying advertised capacity limits. The container exposes it on `GET /healthz` for direct local access and as the response to `POST {"request_type": "health"}` on the unstructured prediction route for DataRobot-hosted deployments. Each endpoint reports only its own capabilities. | | `InstanceHealth` | One configured deployment's probe outcome: `deployment_id`, optional `metadata` (`HealthMetadata` when reachable), and optional `error` when the peer was skipped. | @@ -218,7 +218,7 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | `outputs.log_prob` | `log_prob` responses only: per-horizon `values` and `nll_values` with their `total`, `mean`, `nll_total`, and `nll_mean` summaries. | | `diagnostics.condition_draws` | Condition responses only: number of draws behind sampled outputs. Batched sample requests report the merged count. | | `diagnostics.interval_estimator` | Condition responses only, and only when more than one column carries an interval: `points` and `effective_sample_size` of the numerical region-probability estimate. | -| `plausibility` | `null` on forecast responses. On condition responses an object with `equality_log_density` (log density of the pinned values) and `region_log_probability` (log probability of the interval region), each `null` when the request carried no condition of that kind. | +| `plausibility` | `null` on forecast responses. On condition responses an object with `equality_log_density` (log density of the pinned values) and `region_log_probability` (log probability of the interval region, conditional on the pinned values), each `null` when the request carried no condition of that kind. | | `errors` | Structured service errors. Non-empty arrays raise typed SDK exceptions. | ### Health Metadata From 0c1df59e4752845ba3414f9428e7f5ed485a83c4 Mon Sep 17 00:00:00 2001 From: Stefan Hackmann Date: Thu, 24 Sep 2026 10:58:19 +0000 Subject: [PATCH 8/8] docs: Restore the plausibility chain wording lost in the v5 merge --- README.md | 2 +- docs/api-reference.md | 6 +++--- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/README.md b/README.md index c795d61..658e355 100644 --- a/README.md +++ b/README.md @@ -207,7 +207,7 @@ print(result.plausibility) Whether a deployment can condition depends on the mounted checkpoint's head. `/healthz` advertises `condition` in `supported_query_modes` and the kinds it answers in `supported_condition_kinds` (empty when the mode is absent). The client checks that advertisement before sending, so a deployment that cannot condition is refused with `UnsupportedServiceContractError` rather than after a paid round trip; `require_condition_support(metadata, condition)` exposes the same check. -A condition response carries a `plausibility` block: `equality_log_density` is the log density the model assigns to the pinned values and `region_log_probability` the log probability it gives the interval region, each `None` when the request carried no condition of that kind. Over several covered positions each is the sum of the per-position values, exact because positions are independent. They separate a confident answer from one conditioned on something the model finds implausible; the service reports them and never refuses on them. `diagnostics.condition_draws` counts the draws behind sampled outputs, and `diagnostics.interval_estimator` reports the numerical accounting (`points`, `effective_sample_size`) when some position bounds more than one column and its region probability had to be estimated; across several estimated positions it describes the one with the smallest effective sample size. +A condition response carries a `plausibility` block: `equality_log_density` is the log density the model assigns to the pinned values and `region_log_probability` the log probability it gives the interval region *given the pinned values at the same position*, so the two add to the plausibility of the whole condition set, each `None` when the request carried no condition of that kind. Over several covered positions each is the sum of the per-position values, exact because positions are independent. They separate a confident answer from one conditioned on something the model finds implausible; the service reports them and never refuses on them. `diagnostics.condition_draws` counts the draws behind sampled outputs, and `diagnostics.interval_estimator` reports the numerical accounting (`points`, `effective_sample_size`) when some position bounds more than one column and its region probability had to be estimated; across several estimated positions it describes the one with the smallest effective sample size. Column descriptors support the server fields `name`, `modality`, `role`, `nullable`, `vocabulary_size`, `level_count`, `mapping`, `lower_bound`, `upper_bound`, `time_value_kind`, `time_value_scale_seconds`, `time_value_use_local_normalized_time`, `time_value_calendar_id`, and `time_value_timezone`. diff --git a/docs/api-reference.md b/docs/api-reference.md index ced93bf..e434d76 100644 --- a/docs/api-reference.md +++ b/docs/api-reference.md @@ -21,7 +21,7 @@ This reference covers the supported public Python surface exported by `jointfm_c | `EqualityCondition` | One column pinned to a finite `value` at the positions `query_time_indices` names (indices into `query_times`, `None` for every position). The column may be named in `requested_columns` like any other, and reads back the pinned value the request supplied. | | `IntervalCondition` | One column confined to `[lower, upper]` at the positions `query_time_indices` names (`None` for every position); `None` leaves a side open, at least one side must be bounded, and `lower < upper`. The column stays readable and the response describes its distribution inside the range. | | `Condition` | Type alias for `EqualityCondition \| IntervalCondition`. Every `condition=` parameter takes one of them or a sequence of them. | -| `ConditionPlausibility` | What the model thinks of the conditions it was given: `equality_log_density` of the pinned values and `region_log_probability` of the interval region, each `None` when the request carried no condition of that kind. Reported by the service and never refused on. | +| `ConditionPlausibility` | What the model thinks of the conditions it was given: `equality_log_density` of the pinned values and `region_log_probability` of the interval region given the pinned values at the same position, each `None` when the request carried no condition of that kind. Reported by the service and never refused on. | | `IntervalEstimator` | Numerical accounting (`points`, `effective_sample_size`) behind a region probability that had to be estimated, which happens when some position bounds more than one column. Across several estimated positions it describes the one with the smallest effective sample size. | | `HealthMetadata` | Typed service-health payload with service status, schema and model versions, checkpoint metadata, device, head, `decoding_strategy`, advertised query modes, `supported_condition_kinds` (empty when the deployment cannot condition), return modes, time-index modes, time-index encoding, `max_sample_count`, and an optional `data_generation` block carrying advertised capacity limits. The container exposes it on `GET /healthz` for direct local access and as the response to `POST {"request_type": "health"}` on the unstructured prediction route for DataRobot-hosted deployments. Each endpoint reports only its own capabilities. | | `InstanceHealth` | One configured deployment's probe outcome: `deployment_id`, optional `metadata` (`HealthMetadata` when reachable), and optional `error` when the peer was skipped. | @@ -218,8 +218,8 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | `diagnostics.seed` | Optional seed used by the service. | | `outputs.log_prob` | `log_prob` responses only: per-horizon `values` and `nll_values` with their `total`, `mean`, `nll_total`, and `nll_mean` summaries. | | `diagnostics.condition_draws` | Condition responses only: number of draws behind sampled outputs. Batched sample requests report the merged count. | -| `diagnostics.interval_estimator` | Condition responses only, and only when more than one column carries an interval: `points` and `effective_sample_size` of the numerical region-probability estimate. | -| `plausibility` | `null` on forecast responses. On condition responses an object with `equality_log_density` (log density of the pinned values) and `region_log_probability` (log probability of the interval region), each `null` when the request carried no condition of that kind. Over several covered positions each is the sum of the per-position values. | +| `diagnostics.interval_estimator` | Condition responses only, and only when some position bounds more than one column: `points` and `effective_sample_size` of the numerical region-probability estimate, describing the estimated position with the smallest effective sample size. | +| `plausibility` | `null` on forecast responses. On condition responses an object with `equality_log_density` (log density of the pinned values) and `region_log_probability` (log probability of the interval region, conditional on the pinned values at the same position), each `null` when the request carried no condition of that kind. Over several covered positions each is the sum of the per-position values. | | `errors` | Structured service errors. Non-empty arrays raise typed SDK exceptions. | ### Health Metadata