File size: 5,308 Bytes
ad91e86
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
"""Public-input belief and candidate selection for the N1/N2 high-level track."""

from __future__ import annotations

import math

import numpy as np

from pack_p4d_hssd_records import pack_split
from readyagent.p4d_belief.data import MODEL_INPUT_KEYS, OPTIONAL_MODEL_INPUT_KEYS, derive_last_state


def pack_public_query(query: dict, schema: dict, candidate_features: dict) -> dict[str, np.ndarray]:
    """Reuse the frozen feature schema without passing evaluator labels to the model."""
    query_time = float(query["input"]["query"]["query_time_s"])
    for field in ("target_history", "observable_context_history"):
        timestamps = [float(row["timestamp_s"]) for row in query["input"].get(field, [])]
        if any(timestamp > query_time for timestamp in timestamps):
            raise ValueError(f"future observation in {field}")
        if timestamps != sorted(timestamps):
            raise ValueError(f"out-of-order observation in {field}")
    record = {
        "record_id": query["query_id"],
        "world_variant": "routine",  # packer metadata only; discarded below
        "input": query["input"],
        "supervision": {"current_state_id": 0, "moved_since_last_positive": False},
    }
    packed = pack_split(
        [record],
        max_history=64,
        category_to_id=schema["category_to_id"],
        state_count=int(schema["state_count_including_unknown"]),
        candidate_features=candidate_features,
    )
    allowed = (*MODEL_INPUT_KEYS, *OPTIONAL_MODEL_INPUT_KEYS, "instance_uuid")
    return {key: packed[key] for key in allowed if key in packed}


def candidate_utility(
    *, arrival_probability: float, new_detection_probability: float,
    distance_m: float, eta_s: float, lambda_time: float, lambda_inspect: float,
) -> float:
    """Equation (12): arrival belief times new detection chance per action cost."""
    denominator = distance_m + lambda_time * eta_s + lambda_inspect
    if not math.isfinite(denominator) or denominator <= 0:
        return -math.inf
    return arrival_probability * new_detection_probability / denominator


def rank_candidates(
    belief: dict[int, float], public_costs: dict[int, float], *, task: str = "n2",
    detection_probabilities: dict[int, float] | None = None,
    speed_mps: float = 1.0, inspection_s: float = 1.0,
    lambda_time: float = 0.05, lambda_inspect: float = 0.25,
) -> list[int]:
    """N1 takes Top-1; fixed-target N2 applies Equation (12) to public viewpoints."""
    if not belief:
        return []
    if task == "n1":
        return [max(belief, key=lambda state: (belief[state], -state))]
    if task != "n2":
        raise ValueError(task)
    if speed_mps <= 0:
        raise ValueError("speed_mps must be positive")
    detection_probabilities = detection_probabilities or {}
    return sorted(
        belief,
        key=lambda state: (
            -candidate_utility(
                arrival_probability=belief[state],
                new_detection_probability=detection_probabilities.get(state, 1.0),
                distance_m=public_costs[state],
                eta_s=public_costs[state] / speed_mps + inspection_s,
                lambda_time=lambda_time,
                lambda_inspect=lambda_inspect,
            ),
            state,
        ),
    )


def load_belief(checkpoint_path, dataset_root):
    import torch
    from readyagent.p4d_belief.data import load_catalog
    from readyagent.p4d_belief.models import ModelConfig
    from readyagent.p4d_belief.training import build_model

    catalog = load_catalog(dataset_root)
    checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=True)
    if checkpoint.get("model_name", "p4d" if "transition_head" in checkpoint else None) != "p4d":
        raise ValueError("checkpoint is not a P4D belief model")
    model = build_model("p4d", catalog, ModelConfig(**checkpoint["model_config"]))
    model.load_state_dict(checkpoint["model"])
    return model.eval(), catalog.schema


def model_input_batch(arrays: dict[str, np.ndarray], schema: dict):
    """Strict public-input allowlist for belief and transition inference."""
    import torch

    last_state = int(derive_last_state(arrays)[0])
    candidates = arrays["candidate_state_ids"][0].tolist()
    last_index = candidates.index(last_state)
    instance = str(arrays["instance_uuid"][0])
    vocabulary = schema.get("instance_uuid_to_id", {})
    if instance not in vocabulary:
        raise ValueError(f"unseen instance identity for this checkpoint: {instance}")
    batch = {
        key: torch.as_tensor(arrays[key])
        for key in (*MODEL_INPUT_KEYS, *OPTIONAL_MODEL_INPUT_KEYS)
        if key in arrays
    }
    batch["target_instance_id"] = torch.tensor([vocabulary[instance]])
    batch["last_state"] = torch.tensor([last_state])
    batch["last_candidate_index"] = torch.tensor([last_index])
    return batch


def predict_public(model, schema: dict, arrays: dict[str, np.ndarray]) -> dict[int, float]:
    batch = model_input_batch(arrays, schema)
    candidates = arrays["candidate_state_ids"][0].tolist()
    import torch

    with torch.inference_mode():
        probabilities = model(batch)["probabilities"][0].cpu().numpy()
    return {int(state): float(probabilities[index]) for index, state in enumerate(candidates) if state >= 0}