fishxinyu/flexcombine-bench / data_processing /select_human_subset.py
fishxinyu's picture
download
raw
5.09 kB
import json
import time
from pathlib import Path
import google.generativeai as genai
CAPTION_FILE = "/mnt/data/xinyuy/datasets/VIDGEN-1M/VidGen_1M_sample_3000.json"
OUTPUT_FILE = "vidgen_humans_subset_gemini.json"
GEMINI_API_KEY = "AIzaSyC9O8mR3Fr2c5cPMLtwlyJytNXXJ2uingY" # or use os.environ["GEMINI_API_KEY"]
genai.configure(api_key=GEMINI_API_KEY)
model = genai.GenerativeModel("gemini-2.5-pro") # fast + cheap; swap to gemini-1.5-pro if needed
with open(CAPTION_FILE) as f:
data = json.load(f)
print(f"Total videos: {len(data)}")
# --------------------------------------------------------------------------
# Prompt
# --------------------------------------------------------------------------
# Each call sends a batch of captions so we minimise API round-trips.
# Gemini returns a JSON list of booleans in the same order as the input.
SYSTEM_PROMPT = """\
You are a video-content classifier. Your task is to decide whether each
video caption describes a scene that contains at least one human being
(only real, not animated or cartoon).
Rules:
- Count any depiction of a person, regardless of age, gender.
- Animated or cartoon human characters doesn't count.
- Anthropomorphised animals (e.g. a cartoon pig wearing clothes and acting
human) do NOT count unless explicitly described as human.
- If the caption is ambiguous, lean toward False.
- If they are shown only partially, it doesn't count (e.g. just hands doesn't count as a human).
You will receive a JSON array of caption strings.
Respond with ONLY a JSON array of booleans (true/false), one per caption,
in the same order. No explanation, no markdown fences — raw JSON only.
Example input:
["A woman jogs through the park.", "A dog chases a ball on the grass."]
Example output:
[true, false]
"""
def classify_batch(captions: list[str]) -> list[bool]:
"""Send a batch of captions to Gemini; return a bool list."""
prompt = SYSTEM_PROMPT + "\n" + json.dumps(captions, ensure_ascii=False)
response = model.generate_content(prompt)
text = response.text.strip()
# Strip accidental markdown fences if the model adds them
if text.startswith("```"):
text = text.split("\n", 1)[1].rsplit("```", 1)[0].strip()
results = json.loads(text)
if len(results) != len(captions):
raise ValueError(f"Got {len(results)} results for {len(captions)} captions")
return [bool(r) for r in results]
print("Prompt and classify_batch() defined.")
# Quick smoke test on a few known examples before the full run
test_captions = [
"A woman sits in the driver's seat of a car, smiling at the camera.",
"A close-up of two glasses of water on a car seat.",
"A cartoon shark lifts weights in a gym.",
"A chef chops vegetables on a wooden cutting board.",
]
test_results = classify_batch(test_captions)
for cap, res in zip(test_captions, test_results):
print(f" {'YES' if res else 'NO ':3s} {cap[:80]}")
# --------------------------------------------------------------------------
# Full classification run — saves output after every batch
# --------------------------------------------------------------------------
BATCH_SIZE = 50 # captions per API call — tune down if you hit token limits
SLEEP_SEC = 1.0 # pause between batches to stay within rate limits
labels: list[bool] = []
human_videos: list[dict] = []
for start in range(0, len(data), BATCH_SIZE):
batch = data[start : start + BATCH_SIZE]
captions = [item["caption"] for item in batch]
# Retry once on transient errors
for attempt in range(2):
try:
batch_labels = classify_batch(captions)
labels.extend(batch_labels)
break
except Exception as e:
if attempt == 0:
print(f" Retrying batch {start}–{start+len(batch)} after error: {e}")
time.sleep(3)
else:
print(f" Batch {start}–{start+len(batch)} FAILED twice; marking all False")
batch_labels = [False] * len(batch)
labels.extend(batch_labels)
# Append human videos from this batch and save immediately
batch_humans = [item for item, has_human in zip(batch, batch_labels) if has_human]
human_videos.extend(batch_humans)
with open(OUTPUT_FILE, "w") as f:
json.dump(human_videos, f, indent=2)
end = start + len(batch)
print(f" Processed {end}/{len(data)} "
f"(+{len(batch_humans)} human in this batch, {len(human_videos)} total saved)", flush=True)
time.sleep(SLEEP_SEC)
print(f"\nDone. Videos with humans: {len(human_videos)} / {len(data)} "
f"({100 * len(human_videos) / len(data):.1f}%)")
# Spot-check a few positives and negatives
non_human = [item for item, has_human in zip(data, labels) if not has_human]
print("=== HUMAN (first 5) ===")
for item in human_videos[:5]:
print(f" {item['vid']}")
print(f" {item['caption'][:150]}")
print()
print("=== NON-HUMAN (first 5) ===")
for item in non_human[:5]:
print(f" {item['vid']}")
print(f" {item['caption'][:150]}")
print()

Xet Storage Details

Size:
5.09 kB
·
Xet hash:
56287f0fc38587e626379dfa91d548dcb2943ebb45085aa21e9e7a383ef98808

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.