CoolFace
Apppublic

roi/EditP23

sourceHugging Faceupdated 1y agoView on Hugging Face
5likes
main.py72 linesDownload Raw Back to src
1import argparse2import sys3from pathlib import Path4from edit_mv import run_editp23, load_z123_pipe5 6def main(args: argparse.Namespace) -> None:7    """8    Sets up and runs the EditP23 process for a single experiment.9    """10    exp_dir = Path(args.exp_dir)11    input_files = {12        "src_path": exp_dir / "src.png",13        "edited_path": exp_dir / "edited.png",14        "src_mv_path": exp_dir / "src_mv.png",15    }16 17    # Pre-run validation to ensure all input files exist18    for name, path in input_files.items():19        if not path.is_file():20            print(f"Error: Input file not found at {path}")21            sys.exit(1)22 23    output_dir = exp_dir / "output"24    output_dir.mkdir(exist_ok=True)25    save_path = output_dir / f"result_tgs_{args.tar_guidance_scale}_nmax_{args.n_max}.png"26 27    print(f"Running edit for experiment: {args.exp_dir}")28    print(f"Saving to: {save_path}")29 30    pipeline = load_z123_pipe(args.device_number)31 32    run_editp23(33        src_condition_path=str(input_files["src_path"]),34        tgt_condition_path=str(input_files["edited_path"]),35        original_mv=str(input_files["src_mv_path"]),36        save_path=str(save_path),37        device_number=args.device_number,38        T_steps=args.T_steps,39        n_max=args.n_max,40        src_guidance_scale=args.src_guidance_scale,41        tar_guidance_scale=args.tar_guidance_scale,42        seed=args.seed,43        pipeline=pipeline,44    )45 46if __name__ == "__main__":47    parser = argparse.ArgumentParser(48        description="""Run EditP23 for 3D object editing.49Paper presets for (tar_guidance_scale, n_max):50- Mild: (5, 31)51- Medium: (6, 41), (12, 42)52- Hard: (21, 39)""",53        formatter_class=argparse.RawTextHelpFormatter54    )55 56    parser.add_argument("--exp_dir", type=str, required=True,57                        help="Path to the experiment directory. Expects src.png, edited.png, and src_mv.png in this directory.")58    parser.add_argument("--device_number", type=int, default=0,59                        help="GPU device number to use.")60    parser.add_argument("--seed", type=int, default=18,61                        help="Random seed for reproducibility.")62    parser.add_argument("--T_steps", type=int, default=50,63                        help="Total number of denoising steps.")64    parser.add_argument("--n_max", type=int, default=31,65                        help="Number of scheduler steps for edit-aware guidance. Increase up to T_steps for more significant edits.")66    parser.add_argument("--src_guidance_scale", type=float, default=3.5,67                        help="CFG scale for the source condition. Can typically remain constant.")68    parser.add_argument("--tar_guidance_scale", type=float, default=5.0,69                        help="CFG scale for the target condition. Increase for more significant edits.")70 71    args = parser.parse_args()72    main(args)