SACFlow / tools /make_sfda_manifest.py
sathiiii's picture
Add tools
93ffd19 verified
Raw
History Blame Contribute Delete
1.39 kB
#!/usr/bin/env python
"""Create an SFDA-safe manifest by dropping labels from specified target splits.
This prevents accidental target-label leakage during adaptation. Validation/test labels are
kept by default so metrics can be computed.
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--input", required=True, help="Input manifest JSON")
ap.add_argument("--output", required=True, help="Output manifest JSON")
ap.add_argument("--drop-splits", nargs="+", default=["target_train"], help="Splits whose labels should be removed")
args = ap.parse_args()
with open(args.input, "r") as f:
man = json.load(f)
for split in args.drop_splits:
for item in man.get(split, []):
item.pop("label", None)
man.setdefault("notes", {})["sfda_safe"] = {
"dropped_label_splits": args.drop_splits,
"warning": "Target adaptation labels removed to avoid SFDA leakage."
}
out = Path(args.output)
out.parent.mkdir(parents=True, exist_ok=True)
with open(out, "w") as f:
json.dump(man, f, indent=2)
print(f"Wrote {out}")
for k, v in man.items():
if isinstance(v, list):
print(f"{k}: {len(v)} items, {sum('label' in x for x in v)} labels")
if __name__ == "__main__":
main()