Buckets:
| import json | |
| import re | |
| import sys | |
| from pathlib import Path | |
| import numpy as np | |
| import pandas as pd | |
| import torch | |
| ROOT = Path("/pfss/mlde/workspaces/mlde_wsp_IAS_SAMMerge/triquang/graphvla") | |
| sys.path.insert(0, str(ROOT / "HAMLET-Isaac-GR00T")) | |
| from qgraph.memory.online import load_online | |
| KINDS = ["pick", "place", "table", "press", "move", "static", "insert", "none"] | |
| ORDINALS = ["-", "first", "second", "third", "fourth", "fifth"] | |
| MARGIN, SPAN = 6.4, 243.2 | |
| DATA = ROOT / "data/robomme16" | |
| INFO = json.loads((DATA / "meta/info.json").read_text()) | |
| TASKS = {json.loads(l)["task_index"]: json.loads(l)["task"] for l in open(DATA / "meta/tasks.jsonl")} | |
| LABELS = None | |
| def labels(): | |
| global LABELS | |
| if LABELS is None: | |
| LABELS = torch.load(ROOT / "cache/oracle_labels_16k7.pt", weights_only=False) | |
| return LABELS | |
| def parquet(i): | |
| return DATA / INFO["data_path"].format(episode_chunk=i // INFO["chunks_size"], episode_index=i) | |
| def video(i, key="image"): | |
| return DATA / INFO["video_path"].format(episode_chunk=i // INFO["chunks_size"], video_key=key, episode_index=i) | |
| def episode_text(i): | |
| df = pd.read_parquet(parquet(i), columns=["grounded_subgoal_online", "task_index"]) | |
| subgoals = [str(np.asarray(s).reshape(-1)[0]) if not isinstance(s, str) else s for s in df["grounded_subgoal_online"]] | |
| return TASKS[int(df["task_index"].iloc[0])], subgoals | |
| def load(i, count=192): | |
| item = torch.load(ROOT / f"cache/online16_s4/ep_{i:04d}.pt", weights_only=False) | |
| ent = torch.load(ROOT / f"cache/entities16_p384_s4/ep_{i:04d}.pt", weights_only=False) | |
| return item, ent | |
| def batch(item, ent, device, text=None, count=192): | |
| frames = item["frames"] | |
| n = len(frames) | |
| position = (ent["position"][:, :count].float() - MARGIN) / SPAN | |
| patch = ent["patch"][:, :count].flatten(2, 3) | |
| text = item["text"] if text is None else text | |
| return { | |
| "node_tokens": item["tokens"][None].to(device, torch.bfloat16), | |
| "node_gripper": item["gripper"][None].to(device), | |
| "node_exec": (frames >= item["exec_start"])[None].to(device), | |
| "node_valid": torch.ones(1, n, dtype=torch.bool, device=device), | |
| "text": text[None].to(device, torch.bfloat16), | |
| "text_mask": torch.ones(1, len(text), dtype=torch.bool, device=device), | |
| "node_entity_position": position[None].to(device), | |
| "node_entity_visible": ent["visible"][:, :count][None].to(device), | |
| "node_entity_patch": patch[None].to(device), | |
| } | |
| def run(model, item, ent, device, text=None): | |
| b = batch(item, ent, device, text) | |
| with torch.autocast("cuda", dtype=torch.bfloat16): | |
| out = model(b) | |
| pointer = out["pointer"][0].float().softmax(-1).cpu() | |
| choice = pointer.argmax(-1) | |
| position = b["node_entity_position"][0].cpu() | |
| shift = out["shift"][0].float().cpu() | |
| count = position.shape[1] | |
| idx = choice.clamp(max=count - 1) | |
| point = (position + shift)[torch.arange(len(idx)), idx] | |
| point[choice >= count] = -1 | |
| events = out["events"][0].float().sigmoid().cpu() | |
| return { | |
| "kind": out["kind"][0].float().softmax(-1).cpu(), | |
| "ordinal": out["ordinal"][0].float().softmax(-1).cpu(), | |
| "pointer": pointer, | |
| "point": point, | |
| "counts": events.cumsum(0), | |
| "position": position, | |
| } | |
| def truth(i, frames): | |
| lab = labels()[i] | |
| f = frames.long() | |
| return {"kind": lab["kind"][f], "ordinal": lab["ordinal"][f], "point": lab["point"][f].float()} | |
| def model(device): | |
| return load_online(ROOT / "runs/online/planner16_s0", device)[0].eval() | |
| ORDINAL_WORD = re.compile(r"\b(first|second|third|fourth|fifth)\b(?= target)") | |
Xet Storage Details
- Size:
- 3.67 kB
- Xet hash:
- 63bdcf9834aef484665fb116fcafe207108cc3ba4e5be6c552d2f9b2a93f51bd
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.