fishxinyu/flexcombine-bench / data_processing /select_subject_subset.py
fishxinyu's picture
download
raw
6.72 kB
import json
import shutil
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_subject_subset_gemini.json"
GEMINI_API_KEY = "AIzaSyC9O8mR3Fr2c5cPMLtwlyJytNXXJ2uingY" # or use os.environ["GEMINI_API_KEY"]
VIDEO_DIR = Path("/mnt/data/xinyuy/datasets/VIDGEN-1M/videos")
DEST_DIR = Path("/home/xinyuy/dataset_processing/flexcombine-bench/datasets/raw/subject")
DEST_DIR.mkdir(parents=True, exist_ok=True)
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 objects 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 with a clear, identifiable main subject —
an animal or a distinct object (real, animated, or cartoon). Humans do not count as subjects.
Rules:
- The subject must be the clear focal point of the scene, not background scenery or an incidental prop.
- Anthropomorphised animals (e.g. a cartoon pig wearing clothes and acting human) count as subjects.
- Abstract scenes (landscapes, weather, diffuse crowds without a focal figure) have no main subject.
- If the subject itself is only partially visible (e.g. a disembodied limb, a blurred silhouette), it does not count.
- There must be exactly one main subject; if two or more distinct subjects share equal focus, return false.
- If there are human hand interactions in the video, return false.
- When in doubt, return false.
You will receive a JSON array of caption strings.
Respond with ONLY a JSON array of objects, one per caption, in the same order.
Each object must have:
"has_subject": boolean — whether a clear main subject is present
"subject": string or null — short noun phrase naming the subject (e.g. "cartoon dog"), null if none
"subject_details": string or null — richer noun phrase drawn from the caption (e.g. "a cartoon dog riding a bicycle"), null if none
No explanation, no markdown fences — raw JSON only.
Example input:
["Waves crash against a rocky shore.", "A cartoon dog rides a bicycle."]
Example output:
[{"has_subject": false, "subject": null, "subject_details": null}, {"has_subject": true, "subject": "cartoon dog", "subject_details": "a cartoon dog riding a bicycle"}]
"""
def classify_batch(captions: list[str]) -> list[dict]:
"""Send a batch of captions to Gemini; return list of {has_subject, subject}."""
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 [{"has_subject": r["has_subject"], "subject": r.get("subject"), "subject_details": r.get("subject_details")} 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['has_subject'] 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[dict] = []
subject_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 = [{"has_subject": False, "subject": None}] * len(batch)
labels.extend(batch_labels)
# Append subject videos from this batch and save immediately
batch_subjects = []
for item, label in zip(batch, batch_labels):
if not label["has_subject"]:
continue
enriched = {**item, "subject": label["subject"], "subject_details": label["subject_details"]}
batch_subjects.append(enriched)
src = VIDEO_DIR / (item["vid"] + ".mp4")
dst = DEST_DIR / (item["vid"] + ".mp4")
dst.parent.mkdir(parents=True, exist_ok=True)
if src.exists() and not dst.exists():
shutil.copy2(src, dst)
subject_videos.extend(batch_subjects)
with open(OUTPUT_FILE, "w") as f:
json.dump(subject_videos, f, indent=2)
end = start + len(batch)
print(f" Processed {end}/{len(data)} "
f"(+{len(batch_subjects)} with subject in this batch, {len(subject_videos)} total saved)", flush=True)
time.sleep(SLEEP_SEC)
print(f"\nDone. Videos with subjects: {len(subject_videos)} / {len(data)} "
f"({100 * len(subject_videos) / len(data):.1f}%)")
# Spot-check a few positives and negatives
non_subject = [item for item, label in zip(data, labels) if not label["has_subject"]]
print("=== WITH SUBJECT (first 5) ===")
for item in subject_videos[:5]:
print(f" {item['vid']}")
print(f" {item['caption'][:150]}")
print()
print("=== WITHOUT SUBJECT (first 5) ===")
for item in non_subject[:5]:
print(f" {item['vid']}")
print(f" {item['caption'][:150]}")
print()

Xet Storage Details

Size:
6.72 kB
·
Xet hash:
a6cf3768eca476b24a90d54de54f95c0dbc998e4b79929a75b73a0d29fae31d8

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