Download hy_dev_gen_disca.py from Cccccz/comparison: direct link, hf CLI and curl.
- Browser
- Download file 1.08 kB
-
https://huggingface.co/Cccccz/comparison/resolve/main/hy_dev_gen_disca.py
- Command line
-
hf download hf://Cccccz/comparison/hy_dev_gen_disca.py
-
curl -L -o hy_dev_gen_disca.py https://huggingface.co/Cccccz/comparison/resolve/main/hy_dev_gen_disca.py
1.08 kB
| #!/usr/bin/env python | |
| """Run HY-WorldPlay-DEV-Predictor's generate_predictor_v4_train_case_eval.py with the | |
| DisCa predictor: the current DEV rollout passes ``anchor_distance`` (an ATC input) to | |
| every predictor, which HYWorldPlayPredictorDisCa.forward does not accept. Wrap its | |
| forward to drop every keyword it does not declare; everything else is the stock DEV script.""" | |
| import os, runpy, sys | |
| DEV = "/local/zoubin/cz/projects/HY-WorldPlay-DEV-Predictor" | |
| sys.path.insert(0, DEV); os.chdir(DEV) | |
| import models # noqa: E402 | |
| _orig = models.HYWorldPlayPredictorDisCa.forward | |
| import inspect # noqa: E402 | |
| _accepted = set(inspect.signature(_orig).parameters) | |
| def forward(self, *args, **kw): | |
| # drop inputs the ATC-era rollout passes that DisCa never took | |
| # (anchor_distance, previous_step_chunk_hidden, ...) | |
| return _orig(self, *args, **{k: v for k, v in kw.items() if k in _accepted}) | |
| models.HYWorldPlayPredictorDisCa.forward = forward | |
| sys.argv[0] = os.path.join(DEV, "tools/generate_predictor_v4_train_case_eval.py") | |
| runpy.run_path(sys.argv[0], run_name="__main__") | |