Spaces:
Running on Zero
Running on Zero
File size: 6,824 Bytes
23a59ea | 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 | # preprocess_dataset.py
import json
import multiprocessing as mp
import os
from pathlib import Path
import torch
from torchvision.io import read_image
import torch.nn.functional as F
import argparse
from task_set import TASK_SET, UNSEEN_TASK_SET
def safe_save_frames(frames: torch.Tensor, out_path: Path) -> bool:
"""
Safely save {"frames": frames} to out_path:
- write to a temporary file
- atomically rename to final path
- delete temp file if anything goes wrong
Returns True on success, False on failure.
"""
tmp_path = out_path.with_suffix(out_path.suffix + ".tmp")
try:
# Ensure parent exists
out_path.parent.mkdir(parents=True, exist_ok=True)
# Write to temp
torch.save({"frames": frames}, tmp_path)
# Atomic rename (works across POSIX filesystems)
os.replace(tmp_path, out_path)
print(f" [OK] Saved shard with {frames.shape[0]} frames to {out_path}")
return True
except Exception as e:
print(f" [WARN] Failed saving shard {out_path}: {e}")
# Clean up any partial temp file
try:
if tmp_path.exists():
tmp_path.unlink()
except Exception as e2:
print(f" [WARN] Failed removing temp file {tmp_path}: {e2}")
return False
def process_task(args_tuple):
task, filedir, outdir, target_size, shard_size = args_tuple
task_out_dir = Path(outdir) / task
index_path = task_out_dir / f"{task}_index.json"
# skip if already done (index file is written last, so its presence means success)
if index_path.exists():
print(f"[{task}] already processed, skipping.")
return
task_out_dir.mkdir(parents=True, exist_ok=True)
shard_frames = [] # list of (N_i, 3, target_size, target_size) uint8
total_frames = 0 # running count of buffered frames
shard_idx = 0
shard_meta = {} # shard filename -> num_frames, written as index at the end
i = 0
while True:
png_path = Path(filedir) / f"{task}-{i}.png"
if not png_path.exists():
break
print(f"[{task}] reading {png_path}")
try:
frames = read_image(str(png_path)) # (3, 224, 224 * num_frames), uint8
except Exception as e:
print(f" [WARN] Skipping {png_path} (read error): {e}")
i += 1
continue
C, H, W_total = frames.shape
if H != 224 or W_total % 224 != 0:
print(f" [WARN] Skipping {png_path}, unexpected shape {frames.shape}")
i += 1
continue
num_frames = W_total // 224
if num_frames == 0:
print(f" [WARN] Skipping {png_path}, no frames detected")
i += 1
continue
# Split horizontally: (num_frames, 3, 224, 224)
frames = frames.view(C, 224, num_frames, 224) # (3, 224, N, 224)
frames = frames.permute(2, 0, 1, 3) # (N, 3, 224, 224)
# Downsample if needed; skip entirely when target matches source to avoid
# any float round-trip artifacts in the saved frames.
if target_size == 224:
frames_u8 = frames.contiguous()
else:
frames_f = frames.to(torch.float32) / 255.0
frames_f = F.interpolate(
frames_f,
size=(target_size, target_size),
mode="bilinear",
align_corners=False,
)
frames_u8 = (frames_f.clamp(0.0, 1.0) * 255.0).to(torch.uint8)
shard_frames.append(frames_u8)
total_frames += frames_u8.shape[0]
# Flush complete shards; maintain running total to avoid re-summing the list
while total_frames >= shard_size:
concat = torch.cat(shard_frames, dim=0)
to_save, remainder = concat[:shard_size], concat[shard_size:]
shard_name = f"{task}_shard{shard_idx:04d}.pt"
out_path = task_out_dir / shard_name
print(f"[{task}] saving shard {shard_idx} with {to_save.shape[0]} frames to {out_path}")
ok = safe_save_frames(to_save, out_path)
if ok:
shard_meta[shard_name] = int(to_save.shape[0])
else:
print(f" [WARN] Continuing after failed save of {out_path} (check disk space/FS).")
shard_frames = [remainder] if remainder.shape[0] > 0 else []
total_frames = int(remainder.shape[0])
shard_idx += 1
i += 1
# Flush remainder at the end
if shard_frames:
concat = torch.cat(shard_frames, dim=0)
shard_name = f"{task}_shard{shard_idx:04d}.pt"
out_path = task_out_dir / shard_name
print(f"[{task}] saving final shard {shard_idx} with {concat.shape[0]} frames to {out_path}")
ok = safe_save_frames(concat, out_path)
if ok:
shard_meta[shard_name] = int(concat.shape[0])
# Write metadata index last — its presence signals that the task is complete
with open(index_path, "w") as f:
json.dump(shard_meta, f, indent=2)
print(f"[{task}] wrote index with {len(shard_meta)} shards to {index_path}")
TASK_SET_PRESETS = {
"trained": list(TASK_SET), # 200 training tasks
"unseen": list(UNSEEN_TASK_SET), # 10 held-out tasks
"all": list(TASK_SET) + list(UNSEEN_TASK_SET), # 210 total
}
def main(args):
Path(args.outdir).mkdir(parents=True, exist_ok=True)
if args.tasks is None:
tasks = TASK_SET_PRESETS[args.task_set]
else:
tasks = list(args.tasks)
task_args = [(task, args.filedir, args.outdir, args.target_size, args.shard_size) for task in tasks]
print(f"Processing {len(tasks)} tasks with {args.num_workers} parallel workers")
ctx = mp.get_context("spawn")
with ctx.Pool(processes=args.num_workers) as pool:
pool.map(process_task, task_args)
if __name__ == "__main__":
p = argparse.ArgumentParser()
p.add_argument("--filedir", type=str, default="./data/val")
p.add_argument("--outdir", type=str, default="./data/val-shards")
p.add_argument("--target_size", type=int, default=224)
p.add_argument("--shard_size", type=int, default=4096)
p.add_argument("--num_workers", type=int, default=16)
p.add_argument("--tasks", type=str, nargs="+", default=None,
help="Explicit task list to preprocess. Overrides --task_set.")
p.add_argument("--task_set", type=str, default="trained",
choices=sorted(TASK_SET_PRESETS),
help="Preset task list (used when --tasks is not given). "
"'trained'=TASK_SET (200), 'unseen'=UNSEEN_TASK_SET (10), "
"'all'=union (210).")
main(p.parse_args())
|