amodal-completion-demo / scripts /patch_amodal_paths.py
AlexMurillo's picture
Make data directory writable with fallback for checkpoint downloads
23739c9 verified
Raw
History Blame Contribute Delete
2.22 kB
import os
from pathlib import Path
p = Path('/app/amodal/main.py')
s = p.read_text()
if 'import os' not in s:
s = 'import os\n' + s
s = s.replace('LISA_SERVER_URL = "http://127.0.0.1:7860/"', 'LISA_SERVER_URL = "http://127.0.0.1:7861/"')
s = s.replace('PROJECT_PATH = "/your/path/here/"', 'PROJECT_PATH = os.environ.get("DATA_ROOT", "/data/amodal") + "/"')
s = s.replace('LISA_OUTPUT_PATH = "/your/path/here/LISAoutput/"', 'LISA_OUTPUT_PATH = os.path.join(os.environ.get("DATA_ROOT", "/data/amodal"), "LISAoutput") + "/"')
# The original code uses stabilityai/stable-diffusion-2-inpainting, which is
# gated for this account/token in the deployed Space. Use a public inpainting
# pipeline so inference does not fail on a hidden Hub 401 during the first run.
s = s.replace('"stabilityai/stable-diffusion-2-inpainting"', '"runwayml/stable-diffusion-inpainting"')
s = s.replace('default="Grounded-Segment-Anything/groundingdino_swint_ogc.pth"', 'default=os.path.join(os.environ.get("DATA_ROOT", "/data/amodal"), "checkpoints", "groundingdino_swint_ogc.pth")')
s = s.replace('default="Grounded-Segment-Anything/sam_vit_h_4b8939.pth"', 'default=os.path.join(os.environ.get("DATA_ROOT", "/data/amodal"), "checkpoints", "sam_vit_h_4b8939.pth")')
s = s.replace('default="InstaOrder/InstaOrder_ckpt/InstaOrder_InstaOrderNet_od.pth.tar"', 'default=os.path.join(os.environ.get("DATA_ROOT", "/data/amodal"), "checkpoints", "InstaOrder_InstaOrderNet_od.pth.tar")')
s = s.replace("pretrained='./recognize-anything/ram_plus_swin_large_14m.pth'", "pretrained=os.path.join(os.environ.get('DATA_ROOT', '/data/amodal'), 'checkpoints', 'ram_plus_swin_large_14m.pth')")
# Avoid trying to load LaMa by default; the current paper code says it is not used.
s = s.replace('parser.add_argument(\'--lama_config_path\', type=str, default="lama/big-lama/config.yaml")', 'parser.add_argument(\'--lama_config_path\', type=str, default=None)')
s = s.replace('parser.add_argument(\'--lama_ckpt_path\', type=str, default="lama/big-lama/models/best.ckpt")', 'parser.add_argument(\'--lama_ckpt_path\', type=str, default=None)')
p.write_text(s)
Path(os.environ.get('DATA_ROOT', '/data/amodal'), 'LISAoutput').mkdir(parents=True, exist_ok=True)