EvolvingNav / code /scripts /pack_p4d_hssd_records.py
fengnian1678's picture
Mirror ZJU4EmbodiedAI/EvolvingNav at dfa6872
ad91e86 verified
Raw History Blame Contribute Delete
14.6 kB
#!/usr/bin/env python3
"""Pack P4D JSONL query records into leakage-safe NumPy training tensors."""
from __future__ import annotations
import argparse
import json
import math
from pathlib import Path
from typing import Any
import numpy as np
MAX_HISTORY = 64
MAX_CONTEXT_HISTORY = 16
def arguments() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--root", type=Path, required=True)
parser.add_argument("--max-history", type=int, default=MAX_HISTORY)
return parser.parse_args()
def load_jsonl(path: Path) -> list[dict[str, Any]]:
with path.open(encoding="utf-8") as handle:
return [json.loads(line) for line in handle if line.strip()]
def pack_split(
records: list[dict[str, Any]],
*,
max_history: int,
category_to_id: dict[str, int],
state_count: int,
candidate_features: dict[str, np.ndarray],
) -> dict[str, np.ndarray]:
count = len(records)
event_type = np.zeros((count, max_history), dtype=np.int8)
event_time_days = np.zeros((count, max_history), dtype=np.float32)
delta_time_log1p = np.zeros((count, max_history), dtype=np.float32)
observed_state_id = np.full((count, max_history), -1, dtype=np.int16)
candidate_state_id = np.full((count, max_history), -1, dtype=np.int16)
# detector_conf, identity_conf, visible_fraction, frustum, unoccluded, neg_strength
evidence_features = np.zeros((count, max_history, 6), dtype=np.float32)
history_mask = np.zeros((count, max_history), dtype=np.bool_)
candidate_state_ids = np.full((count, state_count), -1, dtype=np.int16)
candidate_mask = np.zeros((count, state_count), dtype=np.bool_)
target_category_id = np.zeros(count, dtype=np.int16)
context_time_days = np.zeros((count, MAX_CONTEXT_HISTORY), dtype=np.float32)
context_category_counts = np.zeros(
(count, MAX_CONTEXT_HISTORY, len(category_to_id)), dtype=np.float32
)
context_observation_features = np.zeros(
(count, MAX_CONTEXT_HISTORY, 4), dtype=np.float32
)
context_mask = np.zeros((count, MAX_CONTEXT_HISTORY), dtype=np.bool_)
query_time_days = np.zeros(count, dtype=np.float32)
query_time_of_day_sin_cos = np.zeros((count, 2), dtype=np.float32)
query_weekday_id = np.zeros(count, dtype=np.int8)
elapsed_since_last_positive_days = np.zeros(count, dtype=np.float32)
meta_world_variant_id = np.zeros(count, dtype=np.int8)
meta_grounding_quality_id = np.zeros(count, dtype=np.int8)
y_current_state = np.zeros(count, dtype=np.int16)
y_moved = np.zeros(count, dtype=np.bool_)
y_transition_occurred = np.zeros(count, dtype=np.bool_)
y_returned_to_last = np.zeros(count, dtype=np.bool_)
record_ids = []
instance_ids = []
world_ids = {"routine": 0, "random": 1, "static": 2}
grounding_ids = {
"exact": 0,
"role_equivalent": 1,
"fallback": 2,
"counterfactual": 3,
"no_transition": 4,
}
for row_index, record in enumerate(records):
model_input = record["input"]
history = model_input["target_history"][-max_history:]
offset = max_history - len(history)
for history_index, event in enumerate(history, start=offset):
history_mask[row_index, history_index] = True
event_time_days[row_index, history_index] = event["timestamp_s"] / 86400.0
delta_time_log1p[row_index, history_index] = math.log1p(
max(0.0, event["delta_t_from_previous_s"])
)
if event["event_type"] == "positive_observation":
event_type[row_index, history_index] = 1
observed_state_id[row_index, history_index] = event["observed_state_id"]
evidence_features[row_index, history_index, :3] = (
event["detector_confidence"],
event["instance_match_confidence"],
event["visible_fraction_estimate"],
)
elif event["event_type"] == "candidate_inspection":
event_type[row_index, history_index] = 2
candidate_state_id[row_index, history_index] = event["candidate_state_id"]
evidence_features[row_index, history_index, 3:] = (
event["surface_in_frustum_fraction"],
event["surface_unoccluded_fraction"],
event["negative_evidence_strength"],
)
else:
raise ValueError(f"unknown event type: {event['event_type']}")
candidates = model_input["candidate_state_ids"]
if len(candidates) > state_count:
raise ValueError(f"candidate overflow: {record['record_id']}")
candidate_state_ids[row_index, : len(candidates)] = candidates
candidate_mask[row_index, : len(candidates)] = True
target = model_input["target"]
target_category_id[row_index] = category_to_id[target["category"]]
query_time_s = model_input["query"]["query_time_s"]
query_time_days[row_index] = query_time_s / 86400.0
phase = 2.0 * math.pi * (query_time_s % 86400.0) / 86400.0
query_time_of_day_sin_cos[row_index] = (math.sin(phase), math.cos(phase))
query_weekday_id[row_index] = int(query_time_s // 86400) % 7
context_history = model_input.get(
"observable_context_history", []
)[-MAX_CONTEXT_HISTORY:]
context_offset = MAX_CONTEXT_HISTORY - len(context_history)
for context_index, context in enumerate(
context_history, start=context_offset
):
context_mask[row_index, context_index] = True
context_time_days[row_index, context_index] = (
context["timestamp_s"] / 86400.0
)
for category, value in context["observed_category_counts"].items():
context_category_counts[
row_index, context_index, category_to_id[category]
] = float(value)
context_observation_features[row_index, context_index] = (
float(context["detected_object_count"]),
float(context["observed_change_count"]),
float(len(context["inspected_candidate_state_ids"])),
float(len(context["observed_region_ids"])),
)
elapsed_since_last_positive_days[row_index] = (
model_input["history_summary"]["elapsed_since_last_positive_s"] / 86400.0
)
meta_world_variant_id[row_index] = world_ids[record["world_variant"]]
grounding = record["supervision"].get(
"last_transition_grounding_quality", "no_transition"
)
meta_grounding_quality_id[row_index] = grounding_ids.get(
grounding, grounding_ids["no_transition"]
)
y_current_state[row_index] = record["supervision"]["current_state_id"]
y_moved[row_index] = record["supervision"]["moved_since_last_positive"]
y_transition_occurred[row_index] = record["supervision"].get(
"transition_occurred_since_last_positive", y_moved[row_index]
)
y_returned_to_last[row_index] = record["supervision"].get(
"returned_to_last_state", False
)
record_ids.append(record["record_id"])
instance_ids.append(target["instance_uuid"])
return {
"event_type": event_type,
"event_time_days": event_time_days,
"delta_time_log1p": delta_time_log1p,
"observed_state_id": observed_state_id,
"candidate_state_id": candidate_state_id,
"evidence_features": evidence_features,
"history_mask": history_mask,
"candidate_state_ids": candidate_state_ids,
"candidate_mask": candidate_mask,
"target_category_id": target_category_id,
"context_time_days": context_time_days,
"context_category_counts": context_category_counts,
"context_observation_features": context_observation_features,
"context_mask": context_mask,
"query_time_days": query_time_days,
"query_time_of_day_sin_cos": query_time_of_day_sin_cos,
"query_weekday_id": query_weekday_id,
"elapsed_since_last_positive_days": elapsed_since_last_positive_days,
"meta_world_variant_id": meta_world_variant_id,
"meta_grounding_quality_id": meta_grounding_quality_id,
"y_current_state": y_current_state,
"y_moved": y_moved,
"y_transition_occurred": y_transition_occurred,
"y_returned_to_last": y_returned_to_last,
"record_id": np.asarray(record_ids),
"instance_uuid": np.asarray(instance_ids),
**candidate_features,
}
def main() -> int:
args = arguments()
root = args.root.resolve()
objects = load_jsonl(root / "object_pool/objects.jsonl")
categories = sorted({row["category_canonical"] for row in objects})
category_to_id = {category: index for index, category in enumerate(categories)}
state_space = json.loads((root / "scene/candidate_states.json").read_text())
state_count = len(state_space["states"])
all_records = {
split: load_jsonl(root / f"records/{split}_queries.jsonl")
for split in ("train", "val", "test")
}
region_categories = sorted(
{row["region_category"] for row in state_space["states"]}
)
receptacle_categories = sorted(
{row["receptacle_category"] for row in state_space["states"]}
)
region_to_id = {value: index for index, value in enumerate(region_categories)}
receptacle_to_id = {
value: index for index, value in enumerate(receptacle_categories)
}
candidate_features = {
"candidate_region_category_id": np.asarray(
[region_to_id[row["region_category"]] for row in state_space["states"]],
dtype=np.int16,
),
"candidate_receptacle_category_id": np.asarray(
[
receptacle_to_id[row["receptacle_category"]]
for row in state_space["states"]
],
dtype=np.int16,
),
"candidate_center_xyz": np.asarray(
[row["state_center"] or [0.0, 0.0, 0.0] for row in state_space["states"]],
dtype=np.float32,
),
"candidate_is_unknown": np.asarray(
[row["region_id"] == "unknown" for row in state_space["states"]],
dtype=np.bool_,
),
}
summaries = {}
output_dir = root / "records/packed"
output_dir.mkdir(parents=True, exist_ok=True)
for split in ("train", "val", "test"):
records = all_records[split]
tensors = pack_split(
records,
max_history=args.max_history,
category_to_id=category_to_id,
state_count=state_count,
candidate_features=candidate_features,
)
numeric = [value for value in tensors.values() if value.dtype.kind in "biufc"]
if not all(np.isfinite(value).all() for value in numeric):
raise ValueError(f"non-finite tensor in {split}")
if not np.all((tensors["y_current_state"] >= 0) & (tensors["y_current_state"] < state_count)):
raise ValueError(f"invalid state label in {split}")
if not np.all(tensors["history_mask"].any(axis=1)):
raise ValueError(f"empty history in {split}")
positive_mask = tensors["event_type"] == 1
if not np.all(positive_mask.any(axis=1)):
raise ValueError(f"missing positive event in {split}")
path = output_dir / f"{split}.npz"
np.savez_compressed(path, **tensors)
summaries[split] = {
"records": len(records),
"path": str(path.relative_to(root)),
"bytes": path.stat().st_size,
"history_length_min": int(tensors["history_mask"].sum(axis=1).min()),
"history_length_max": int(tensors["history_mask"].sum(axis=1).max()),
"moved_fraction": float(tensors["y_moved"].mean()),
"transition_occurred_fraction": float(
tensors["y_transition_occurred"].mean()
),
"returned_to_last_fraction": float(
tensors["y_returned_to_last"].mean()
),
}
schema = {
"schema_version": "p4d_packed_tensor_v1.2",
"max_history": args.max_history,
"state_count_including_unknown": state_count,
"category_to_id": category_to_id,
"region_category_to_id": region_to_id,
"receptacle_category_to_id": receptacle_to_id,
"metadata_world_variant_to_id": {"routine": 0, "random": 1, "static": 2},
"metadata_grounding_quality_to_id": {
"exact": 0,
"role_equivalent": 1,
"fallback": 2,
"counterfactual": 3,
"no_transition": 4,
},
"event_type_to_id": {"padding": 0, "positive_observation": 1, "candidate_inspection": 2},
"evidence_feature_order": [
"detector_confidence",
"instance_match_confidence",
"visible_fraction_estimate",
"surface_in_frustum_fraction",
"surface_unoccluded_fraction",
"negative_evidence_strength",
],
"observable_context_feature_order": [
"detected_object_count",
"observed_change_count",
"inspected_candidate_state_count",
"observed_region_count",
],
"label_tensors": [
"y_current_state",
"y_moved",
"y_transition_occurred",
"y_returned_to_last",
],
"metadata_tensors_not_for_model_input": [
"meta_world_variant_id",
"meta_grounding_quality_id",
"record_id",
"instance_uuid",
],
"access_rule": "input tensors derive from query.input; y_* derives from supervision; meta_* is only for filtering and reporting",
"pytorch_dtype_note": "cast int16 state/category ID arrays to torch.long before embedding, one_hot, or cross_entropy",
"splits": summaries,
"validation": {
"finite_numeric_tensors": True,
"valid_label_ranges": True,
"nonempty_history": True,
"positive_anchor_present": True,
},
}
(output_dir / "feature_schema.json").write_text(
json.dumps(schema, ensure_ascii=False, indent=2), encoding="utf-8"
)
print(json.dumps(schema, ensure_ascii=False, indent=2))
return 0
if __name__ == "__main__":
raise SystemExit(main())