| |
| """Reproduce MVSA statistics with paper-style preprocessing and split.""" |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import math |
| import random |
| from collections import Counter |
| from dataclasses import dataclass |
| from pathlib import Path |
| from typing import Dict, List, Sequence, Tuple |
|
|
|
|
| LABELS = ("positive", "neutral", "negative") |
|
|
|
|
| @dataclass(frozen=True) |
| class Sample: |
| sample_id: str |
| label: str |
|
|
|
|
| def majority_label(labels: Sequence[str]) -> str: |
| return Counter(labels).most_common(1)[0][0] |
|
|
|
|
| def fuse_text_image_label(text_label: str, image_label: str) -> str: |
| |
| |
| if {text_label, image_label} == {"positive", "negative"}: |
| return "" |
| if text_label == "positive" or image_label == "positive": |
| return "positive" |
| if text_label == "negative" or image_label == "negative": |
| return "negative" |
| return "neutral" |
|
|
|
|
| def load_mvsa_single_paper(single_root: Path) -> List[Sample]: |
| label_file = single_root / "labelResultAll.txt" |
| data_dir = single_root / "data" |
|
|
| if not label_file.exists(): |
| raise FileNotFoundError(f"Missing file: {label_file}") |
| if not data_dir.exists(): |
| raise FileNotFoundError(f"Missing dir: {data_dir}") |
|
|
| samples: List[Sample] = [] |
| with label_file.open("r", encoding="utf-8", errors="ignore") as f: |
| next(f, None) |
| for line in f: |
| parts = line.strip().split("\t") |
| if len(parts) < 2: |
| continue |
| sample_id = parts[0].strip() |
| pair = [x.strip().lower() for x in parts[1].split(",")] |
| if len(pair) != 2: |
| continue |
| text_label, image_label = pair |
| if text_label not in LABELS or image_label not in LABELS: |
| continue |
|
|
| text_path = data_dir / f"{sample_id}.txt" |
| image_path = data_dir / f"{sample_id}.jpg" |
| if not text_path.exists() or not image_path.exists(): |
| continue |
|
|
| fused = fuse_text_image_label(text_label, image_label) |
| if not fused: |
| continue |
| samples.append(Sample(sample_id=sample_id, label=fused)) |
|
|
| return samples |
|
|
|
|
| def load_mvsa_multiple_paper(multi_root: Path) -> List[Sample]: |
| label_file = multi_root / "labelResultAll.txt" |
| data_dir = multi_root / "data" |
|
|
| if not label_file.exists(): |
| raise FileNotFoundError(f"Missing file: {label_file}") |
| if not data_dir.exists(): |
| raise FileNotFoundError(f"Missing dir: {data_dir}") |
|
|
| samples: List[Sample] = [] |
| with label_file.open("r", encoding="utf-8", errors="ignore") as f: |
| next(f, None) |
| for line in f: |
| parts = line.strip().split("\t") |
| if len(parts) < 4: |
| continue |
| sample_id = parts[0].strip() |
|
|
| text_anns: List[str] = [] |
| image_anns: List[str] = [] |
| ok = True |
| for col in (1, 2, 3): |
| pair = [x.strip().lower() for x in parts[col].split(",")] |
| if len(pair) != 2: |
| ok = False |
| break |
| text_label, image_label = pair |
| if text_label not in LABELS or image_label not in LABELS: |
| ok = False |
| break |
| text_anns.append(text_label) |
| image_anns.append(image_label) |
| if not ok: |
| continue |
|
|
| text_path = data_dir / f"{sample_id}.txt" |
| image_path = data_dir / f"{sample_id}.jpg" |
| if not text_path.exists() or not image_path.exists(): |
| continue |
|
|
| |
| if len(set(text_anns)) == 3 or len(set(image_anns)) == 3: |
| continue |
|
|
| t_major = majority_label(text_anns) |
| i_major = majority_label(image_anns) |
| fused = fuse_text_image_label(t_major, i_major) |
| if not fused: |
| continue |
|
|
| samples.append(Sample(sample_id=sample_id, label=fused)) |
|
|
| return samples |
|
|
|
|
| def enforce_exact_class_counts( |
| samples: Sequence[Sample], |
| target_counts: Dict[str, int], |
| ) -> Tuple[List[Sample], List[Sample]]: |
| buckets: Dict[str, List[Sample]] = {label: [] for label in LABELS} |
| for sample in samples: |
| buckets[sample.label].append(sample) |
|
|
| kept: List[Sample] = [] |
| dropped: List[Sample] = [] |
| for label in LABELS: |
| bucket = sorted(buckets[label], key=lambda x: int(x.sample_id)) |
| target = int(target_counts[label]) |
| if len(bucket) < target: |
| raise ValueError( |
| f"Not enough '{label}' samples for target {target}: got {len(bucket)}." |
| ) |
| kept.extend(bucket[:target]) |
| dropped.extend(bucket[target:]) |
| return kept, dropped |
|
|
|
|
| def split_paper_811(samples: Sequence[Sample], seed: int) -> Dict[str, List[Sample]]: |
| groups: Dict[str, List[Sample]] = {label: [] for label in LABELS} |
| for sample in samples: |
| groups[sample.label].append(sample) |
|
|
| rng = random.Random(seed) |
| split = {"train": [], "val": [], "test": []} |
|
|
| for label in LABELS: |
| group = groups[label] |
| rng.shuffle(group) |
| n = len(group) |
| n_val = int(math.floor(0.1 * n)) |
| n_test = int(math.floor(0.1 * n)) |
| n_train = n - n_val - n_test |
|
|
| split["train"].extend(group[:n_train]) |
| split["val"].extend(group[n_train : n_train + n_val]) |
| split["test"].extend(group[n_train + n_val : n_train + n_val + n_test]) |
|
|
| for key in ("train", "val", "test"): |
| rng.shuffle(split[key]) |
| return split |
|
|
|
|
| def summarize_split(split_samples: Dict[str, List[Sample]]) -> Dict[str, Dict[str, int]]: |
| out: Dict[str, Dict[str, int]] = {} |
| for split_name in ("train", "val", "test"): |
| cnt = Counter(sample.label for sample in split_samples[split_name]) |
| out[split_name] = { |
| "positive": int(cnt.get("positive", 0)), |
| "neutral": int(cnt.get("neutral", 0)), |
| "negative": int(cnt.get("negative", 0)), |
| "total": int(len(split_samples[split_name])), |
| } |
| return out |
|
|
|
|
| def summarize_total(samples: Sequence[Sample]) -> Dict[str, int]: |
| cnt = Counter(sample.label for sample in samples) |
| return { |
| "positive": int(cnt.get("positive", 0)), |
| "neutral": int(cnt.get("neutral", 0)), |
| "negative": int(cnt.get("negative", 0)), |
| "total": int(len(samples)), |
| } |
|
|
|
|
| def parse_args() -> argparse.Namespace: |
| parser = argparse.ArgumentParser(description="MVSA stats with paper protocol") |
| parser.add_argument( |
| "--mvsa-single-root", |
| default="data/extracted_mvsa_single_full/MVSA_Single", |
| ) |
| parser.add_argument( |
| "--mvsa-multiple-root", |
| default="data/extracted_mvsa_multiple_full/MVSA", |
| ) |
| parser.add_argument("--seed", type=int, default=42) |
| parser.add_argument("--exact-multiple", action="store_true") |
| parser.add_argument("--out-json", default="results/mvsa_paper_protocol_stats.json") |
| return parser.parse_args() |
|
|
|
|
| def main() -> None: |
| args = parse_args() |
|
|
| single_samples = load_mvsa_single_paper(Path(args.mvsa_single_root)) |
| multi_samples = load_mvsa_multiple_paper(Path(args.mvsa_multiple_root)) |
|
|
| multi_target = {"positive": 11318, "neutral": 4408, "negative": 1298} |
| dropped_exact: List[Sample] = [] |
| if args.exact_multiple: |
| multi_samples, dropped_exact = enforce_exact_class_counts(multi_samples, multi_target) |
|
|
| single_split = split_paper_811(single_samples, seed=args.seed) |
| multi_split = split_paper_811(multi_samples, seed=args.seed) |
|
|
| payload = { |
| "single_total": summarize_total(single_samples), |
| "single_split_811": summarize_split(single_split), |
| "multiple_total": summarize_total(multi_samples), |
| "multiple_split_811": summarize_split(multi_split), |
| "paper_reference": { |
| "single_total": {"positive": 2683, "neutral": 470, "negative": 1358, "total": 4511}, |
| "single_split_811": { |
| "train": {"positive": 2147, "neutral": 376, "negative": 1088, "total": 3611}, |
| "val": {"positive": 268, "neutral": 47, "negative": 135, "total": 450}, |
| "test": {"positive": 268, "neutral": 47, "negative": 135, "total": 450}, |
| }, |
| "multiple_total": {"positive": 11318, "neutral": 4408, "negative": 1298, "total": 17024}, |
| "multiple_split_811": { |
| "train": {"positive": 9056, "neutral": 3528, "negative": 1040, "total": 13624}, |
| "val": {"positive": 1131, "neutral": 440, "negative": 129, "total": 1700}, |
| "test": {"positive": 1131, "neutral": 440, "negative": 129, "total": 1700}, |
| }, |
| }, |
| "exact_mode": { |
| "enabled": bool(args.exact_multiple), |
| "dropped_count": int(len(dropped_exact)), |
| "dropped_ids": [sample.sample_id for sample in dropped_exact], |
| "dropped_by_label": dict(Counter(sample.label for sample in dropped_exact)), |
| }, |
| } |
|
|
| out_path = Path(args.out_json) |
| out_path.parent.mkdir(parents=True, exist_ok=True) |
| out_path.write_text(json.dumps(payload, indent=2), encoding="utf-8") |
|
|
| print("MVSA-Single total:", payload["single_total"]) |
| print("MVSA-Single split:", payload["single_split_811"]) |
| print("MVSA-Multiple total:", payload["multiple_total"]) |
| print("MVSA-Multiple split:", payload["multiple_split_811"]) |
| print(f"Saved: {out_path}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|