Instructions to use phi-lab-rice/GRADE with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use phi-lab-rice/GRADE with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("phi-lab-rice/GRADE", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
File size: 8,509 Bytes
f348660 | 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 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 | #!/usr/bin/env python3
"""
Inference Script for GRT-Small (finetuned weights)
Runs inference on all valid sequences in the configured Smoke-Eval root by
default, or on an explicit list supplied with ``--sequences``.
using weights trained by grt_finetune/train.py.
For each sequence, saves one .npy file: pred_depth.npy (dequantized predicted depth [T, 64, 128], values in [0, 1]).
Single GPU: Each frame is seen exactly once; no duplication or incompleteness.
Multi-GPU (DDP): Dataloader is sharded; each rank writes its results to a file, then
main process merges with deduplication by frame_idx (keeps first occurrence) and saves.
"""
import os
import torch
import numpy as np
import argparse
import yaml
import pickle
from tqdm import tqdm
from accelerate import Accelerator
from accelerate.utils import set_seed
from collections import defaultdict
from safetensors.torch import load_file
from grt_model import GRTSmall
from dataloader import create_rice_dataloader
from augmentations import (
translate_radar,
dequantize_depth,
)
def batch_radar_to_spectrum(
radar_amplitude: torch.Tensor, radar_phase: torch.Tensor
) -> torch.Tensor:
"""Restore the GRT spectrum layout from the packaged Smoke-Eval tensors."""
amplitude = radar_amplitude.permute(0, 1, 3, 2, 4)
phase = radar_phase.permute(0, 1, 3, 2, 4)
return torch.stack((amplitude, phase), dim=-1)
def main():
parser = argparse.ArgumentParser(
description="Run GRT inference on Smoke-Eval."
)
parser.add_argument(
"--config", type=str, default="config.yaml", help="Path to config file"
)
parser.add_argument(
"--checkpoint",
type=str,
required=True,
help="Path to weights-only GRT .safetensors file",
)
parser.add_argument(
"--output_dir",
type=str,
default="inference_results",
help="Directory to save results",
)
parser.add_argument(
"--sequences",
type=str,
nargs="+",
default=None,
help="Optional sequence names; default discovers all valid sequences.",
)
parser.add_argument(
"--debug", action="store_true", help="Run in debug mode (process only 1 batch)"
)
args = parser.parse_args()
# Load config
with open(args.config, "r") as f:
config = yaml.safe_load(f)
# Initialize accelerator
accelerator = Accelerator(mixed_precision="fp16")
set_seed(config["training"].get("seed", 42))
# Create output directory (all ranks so DDP gather_dir can be created)
os.makedirs(args.output_dir, exist_ok=True)
# Create model
accelerator.print("Creating GRT-Small model...")
model = GRTSmall()
# Safetensors files contain only the model state dictionary.
accelerator.print(f"Loading checkpoint from {args.checkpoint}")
model.load_state_dict(load_file(args.checkpoint, device="cpu"), strict=True)
# With ``sequences=None`` the public dataset loader discovers every valid
# sequence under the configured Smoke-Eval root.
accelerator.print(f"Inference sequences: {args.sequences}")
inference_loader = create_rice_dataloader(
root_dir=config["paths"]["data_root"],
batch_size=config["training"]["batch_size"],
num_workers=0,
frame_skip=1,
sequences=args.sequences,
shuffle=False,
)
# Prepare model and dataloader
model, inference_loader = accelerator.prepare(model, inference_loader)
model.eval()
# Dictionary to aggregate results by sequence: sequence_id -> list of (frame_idx, pred_depth)
results_by_sequence = defaultdict(list)
accelerator.print("Starting inference...")
with torch.no_grad():
for batch in tqdm(
inference_loader, disable=not accelerator.is_local_main_process
):
# Extract data
rsp_data = batch_radar_to_spectrum(
batch["radar_amplitude"], batch["radar_phase"]
)
sequences = batch["sequence"]
frame_indices = batch["frame_idx"]
# Apply radar augmentation
rsp_data = translate_radar(rsp_data)
# Forward pass
occupancy_pred_logits = model(rsp_data) # [B, 64, 128, 64]
# Dequantize predicted occupancy to depth [B, 1, 64, 128], values in [0, 1]
pred_depth = dequantize_depth(occupancy_pred_logits)
pred_depth_np = (
pred_depth.cpu().numpy().astype(np.float32)
) # [B, 1, 64, 128]
# Collect results (frame_idx, pred_depth per sample)
for i in range(len(sequences)):
seq_id = sequences[i]
f_idx = frame_indices[i].item()
# Store [1, 64, 128] per frame; will stack to [T, 64, 128] when saving
results_by_sequence[seq_id].append(
{
"frame_idx": f_idx,
"pred_depth": pred_depth_np[i],
}
)
if args.debug:
break
# Single GPU: save directly (each frame seen once, no duplication)
# Multi-GPU: gather via files, merge with dedupe by frame_idx, then save
if accelerator.num_processes == 1:
if accelerator.is_main_process:
accelerator.print("Saving results (single process)...")
for seq_id, frames in tqdm(
results_by_sequence.items(), desc="Saving sequences"
):
frames.sort(key=lambda x: x["frame_idx"])
pred_depth_stack = np.stack([f["pred_depth"] for f in frames], axis=0)
pred_depth_stack = np.squeeze(pred_depth_stack, axis=1) # [T, 64, 128]
np.save(
os.path.join(args.output_dir, f"{seq_id.lower()}_pred.npy"),
pred_depth_stack,
)
accelerator.print(
f" {seq_id}: saved {pred_depth_stack.shape[0]} frames"
)
accelerator.print(f"Processed {len(results_by_sequence)} sequences.")
accelerator.print(f"Results saved to {args.output_dir}")
else:
# DDP: gather results from all ranks via files, dedupe by frame_idx, save on main
accelerator.wait_for_everyone()
gather_dir = os.path.join(args.output_dir, "_gather")
os.makedirs(gather_dir, exist_ok=True)
rank = accelerator.process_index
rank_file = os.path.join(gather_dir, f"rank_{rank}_results.pkl")
with open(rank_file, "wb") as f:
pickle.dump(dict(results_by_sequence), f, protocol=pickle.HIGHEST_PROTOCOL)
accelerator.wait_for_everyone()
if accelerator.is_main_process:
accelerator.print("Merging and deduplicating results from all ranks...")
merged_results = defaultdict(dict) # seq_id -> {frame_idx: pred_depth}
for r in range(accelerator.num_processes):
pkl_path = os.path.join(gather_dir, f"rank_{r}_results.pkl")
with open(pkl_path, "rb") as f:
rank_results = pickle.load(f)
for seq_id, frames in rank_results.items():
for frame_data in frames:
f_idx = frame_data["frame_idx"]
if f_idx not in merged_results[seq_id]:
merged_results[seq_id][f_idx] = frame_data["pred_depth"]
os.remove(pkl_path)
for seq_id, frame_dict in tqdm(
merged_results.items(), desc="Saving sequences"
):
sorted_items = sorted(frame_dict.items(), key=lambda x: x[0])
pred_depth_stack = np.stack([item[1] for item in sorted_items], axis=0)
pred_depth_stack = np.squeeze(pred_depth_stack, axis=1) # [T, 64, 128]
np.save(
os.path.join(args.output_dir, f"{seq_id.lower()}_pred.npy"),
pred_depth_stack,
)
accelerator.print(
f" {seq_id}: saved {pred_depth_stack.shape[0]} frames"
)
if os.path.isdir(gather_dir) and not os.listdir(gather_dir):
os.rmdir(gather_dir)
accelerator.print(f"Processed {len(merged_results)} sequences.")
accelerator.print(f"Results saved to {args.output_dir}")
accelerator.wait_for_everyone()
if __name__ == "__main__":
main()
|