twanghcmut's picture
download
raw
3.67 kB
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),
}
@torch.no_grad()
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.