diff --git a/.gitignore b/.gitignore index 1cda7444c..16e4a1210 100644 --- a/.gitignore +++ b/.gitignore @@ -17,9 +17,10 @@ data_* *.gv *.png *.csv +*.jsonl +*.json pkg/ -src/ ### CVS template /CVS/* diff --git a/jericho_priorzero_1020.yml b/jericho_priorzero_1020.yml deleted file mode 100644 index f4ef0b339..000000000 --- a/jericho_priorzero_1020.yml +++ /dev/null @@ -1,454 +0,0 @@ -name: base -channels: - - pytorch - - nvidia - - defaults -dependencies: - - _libgcc_mutex=0.1=main - - _openmp_mutex=5.1=1_gnu - - anaconda-anon-usage=0.4.4=py310hc06175d_0 - - archspec=0.2.3=pyhd3eb1b0_0 - - asttokens=2.0.5=pyhd3eb1b0_0 - - attrs=23.1.0=py310h06a4308_0 - - beautifulsoup4=4.12.2=py310h06a4308_0 - - blas=1.0=mkl - - boltons=23.0.0=py310h06a4308_0 - - brotli-python=1.0.9=py310h6a678d5_8 - - bzip2=1.0.8=h5eee18b_6 - - c-ares=1.19.1=h5eee18b_0 - - ca-certificates=2024.3.11=h06a4308_0 - - certifi=2024.2.2=py310h06a4308_0 - - cffi=1.16.0=py310h5eee18b_1 - - chardet=4.0.0=py310h06a4308_1003 - - charset-normalizer=2.0.4=pyhd3eb1b0_0 - - click=8.1.7=py310h06a4308_0 - - cmake=3.26.4=h96355d8_0 - - conda=23.5.2=py310h06a4308_0 - - conda-build=24.3.0=py310h06a4308_0 - - conda-content-trust=0.2.0=py310h06a4308_1 - - conda-index=0.4.0=pyhd3eb1b0_0 - - conda-libmamba-solver=23.7.0=py310h06a4308_0 - - conda-package-handling=2.2.0=py310h06a4308_1 - - conda-package-streaming=0.9.0=py310h06a4308_0 - - cryptography=42.0.5=py310hdda0065_1 - - cuda-cudart=12.1.105=0 - - cuda-cupti=12.1.105=0 - - cuda-libraries=12.1.0=0 - - cuda-nvrtc=12.1.105=0 - - cuda-nvtx=12.1.105=0 - - cuda-opencl=12.5.39=0 - - cuda-runtime=12.1.0=0 - - cuda-version=12.5=3 - - distro=1.9.0=py310h06a4308_0 - - exceptiongroup=1.2.0=py310h06a4308_0 - - executing=0.8.3=pyhd3eb1b0_0 - - expat=2.6.2=h6a678d5_0 - - ffmpeg=4.3=hf484d3e_0 - - fmt=9.1.0=hdb19cb5_1 - - freetype=2.12.1=h4a9f257_0 - - frozendict=2.4.2=py310h5eee18b_0 - - gmp=6.2.1=h295c915_3 - - gmpy2=2.1.2=py310heeb90bb_0 - - gnutls=3.6.15=he1e5248_0 - - icu=73.1=h6a678d5_0 - - idna=3.7=py310h06a4308_0 - - intel-openmp=2023.1.0=hdb19cb5_46306 - - ipython=8.20.0=py310h06a4308_0 - - jedi=0.18.1=py310h06a4308_1 - - jpeg=9e=h5eee18b_1 - - jsonpatch=1.33=py310h06a4308_1 - - jsonpointer=2.1=pyhd3eb1b0_0 - - jsonschema-specifications=2023.7.1=py310h06a4308_0 - - krb5=1.20.1=h143b758_1 - - lame=3.100=h7b6447c_0 - - lcms2=2.12=h3be6417_0 - - ld_impl_linux-64=2.38=h1181459_1 - - lerc=3.0=h295c915_0 - - libarchive=3.6.2=h6ac8c49_3 - - libcublas=12.1.0.26=0 - - libcufft=11.0.2.4=0 - - libcufile=1.10.0.4=0 - - libcurand=10.3.6.39=0 - - libcurl=8.7.1=h251f7ec_0 - - libcusolver=11.4.4.55=0 - - libcusparse=12.0.2.55=0 - - libdeflate=1.17=h5eee18b_1 - - libedit=3.1.20230828=h5eee18b_0 - - libev=4.33=h7f8727e_1 - - libffi=3.4.4=h6a678d5_1 - - libgcc-ng=11.2.0=h1234567_1 - - libgomp=11.2.0=h1234567_1 - - libiconv=1.16=h5eee18b_3 - - libidn2=2.3.4=h5eee18b_0 - - libjpeg-turbo=2.0.0=h9bf148f_0 - - liblief=0.12.3=h6a678d5_0 - - libmamba=1.5.8=hfe524e5_2 - - libmambapy=1.5.8=py310h2dafd23_2 - - libnghttp2=1.57.0=h2d74bed_0 - - libnpp=12.0.2.50=0 - - libnvjitlink=12.1.105=0 - - libnvjpeg=12.1.1.14=0 - - libpng=1.6.39=h5eee18b_0 - - libsolv=0.7.24=he621ea3_1 - - libssh2=1.11.0=h251f7ec_0 - - libstdcxx-ng=11.2.0=h1234567_1 - - libtasn1=4.19.0=h5eee18b_0 - - libtiff=4.5.1=h6a678d5_0 - - libunistring=0.9.10=h27cfd23_0 - - libuuid=1.41.5=h5eee18b_0 - - libuv=1.44.2=h5eee18b_0 - - libwebp-base=1.3.2=h5eee18b_0 - - libxml2=2.10.4=hfdd30dd_2 - - llvm-openmp=14.0.6=h9e868ea_0 - - lz4-c=1.9.4=h6a678d5_1 - - markupsafe=2.1.3=py310h5eee18b_0 - - matplotlib-inline=0.1.6=py310h06a4308_0 - - menuinst=2.1.0=py310h06a4308_0 - - mkl=2023.1.0=h213fc3f_46344 - - mkl-service=2.4.0=py310h5eee18b_1 - - mkl_fft=1.3.8=py310h5eee18b_0 - - mkl_random=1.2.4=py310hdb19cb5_0 - - more-itertools=10.1.0=py310h06a4308_0 - - mpc=1.1.0=h10f8cd9_1 - - mpfr=4.0.2=hb69a4c5_1 - - mpmath=1.3.0=py310h06a4308_0 - - ncurses=6.4=h6a678d5_0 - - nettle=3.7.3=hbbd107a_1 - - numpy=1.26.4=py310h5f9d8c6_0 - - numpy-base=1.26.4=py310hb5e798b_0 - - openh264=2.1.1=h4ff587b_0 - - openjpeg=2.4.0=h3ad879b_0 - - openssl=3.0.13=h7f8727e_2 - - packaging=23.2=py310h06a4308_0 - - parso=0.8.3=pyhd3eb1b0_0 - - patch=2.7.6=h7b6447c_1001 - - patchelf=0.17.2=h6a678d5_0 - - pcre2=10.42=hebb0a14_1 - - pexpect=4.8.0=pyhd3eb1b0_3 - - pillow=10.3.0=py310h5eee18b_0 - - pkginfo=1.10.0=py310h06a4308_0 - - platformdirs=3.10.0=py310h06a4308_0 - - prompt-toolkit=3.0.43=py310h06a4308_0 - - prompt_toolkit=3.0.43=hd3eb1b0_0 - - psutil=5.9.0=py310h5eee18b_0 - - ptyprocess=0.7.0=pyhd3eb1b0_2 - - pure_eval=0.2.2=pyhd3eb1b0_0 - - py-lief=0.12.3=py310h6a678d5_0 - - pybind11-abi=4=hd3eb1b0_1 - - pycosat=0.6.6=py310h5eee18b_1 - - pycparser=2.21=pyhd3eb1b0_0 - - pygments=2.15.1=py310h06a4308_1 - - pyopenssl=24.0.0=py310h06a4308_0 - - pysocks=1.7.1=py310h06a4308_0 - - python=3.10.14=h955ad1f_1 - - python-libarchive-c=2.9=pyhd3eb1b0_1 - - pytorch-cuda=12.1=ha16c6d3_5 - - pytorch-mutex=1.0=cuda - - pytz=2024.1=py310h06a4308_0 - - pyyaml=6.0.1=py310h5eee18b_0 - - readline=8.2=h5eee18b_0 - - referencing=0.30.2=py310h06a4308_0 - - reproc=14.2.4=h6a678d5_2 - - reproc-cpp=14.2.4=h6a678d5_2 - - requests=2.32.2=py310h06a4308_0 - - rhash=1.4.3=hdbd6064_0 - - rpds-py=0.10.6=py310hb02cf49_0 - - ruamel.yaml=0.17.21=py310h5eee18b_0 - - ruamel.yaml.clib=0.2.6=py310h5eee18b_1 - - six=1.16.0=pyhd3eb1b0_1 - - soupsieve=2.5=py310h06a4308_0 - - sqlite=3.45.3=h5eee18b_0 - - stack_data=0.2.0=pyhd3eb1b0_0 - - tbb=2021.8.0=hdb19cb5_0 - - tk=8.6.14=h39e8969_0 - - tomli=2.0.1=py310h06a4308_0 - - toolz=0.12.0=py310h06a4308_0 - - tqdm=4.66.4=py310h2f386ee_0 - - traitlets=5.7.1=py310h06a4308_0 - - truststore=0.8.0=py310h06a4308_0 - - urllib3=2.2.1=py310h06a4308_0 - - wcwidth=0.2.5=pyhd3eb1b0_0 - - wheel=0.43.0=py310h06a4308_0 - - xz=5.4.6=h5eee18b_1 - - yaml=0.2.5=h7b6447c_0 - - yaml-cpp=0.8.0=h6a678d5_1 - - zlib=1.2.13=h5eee18b_1 - - zstandard=0.22.0=py310h2c38b39_0 - - zstd=1.5.5=hc292b87_2 - - pip: - - absl-py==2.1.0 - - accelerate==1.10.1 - - aiohappyeyeballs==2.4.0 - - aiohttp==3.10.5 - - aiosignal==1.3.1 - - ale-py==0.8.1 - - annotated-types==0.7.0 - - anyio==4.11.0 - - astor==0.8.1 - - astunparse==1.6.3 - - async-timeout==4.0.3 - - av==12.3.0 - - beartype==0.18.5 - - bitmath==1.3.3.1 - - blake3==1.0.8 - - blis==1.3.0 - - box2d-py==2.3.5 - - cachetools==6.2.1 - - catalogue==2.0.10 - - cbor2==5.7.0 - - cloudpathlib==0.23.0 - - cloudpickle==3.0.0 - - comm==0.2.2 - - compressed-tensors==0.11.0 - - confection==0.1.5 - - contourpy==1.2.1 - - cupy-cuda12x==13.6.0 - - cycler==0.12.1 - - cymem==2.0.11 - - cython==0.29.37 - - datasets==4.2.0 - - debugpy==1.8.5 - - decorator==4.4.2 - - deprecation==2.1.0 - - depyf==0.19.0 - - di-engine==0.5.3 - - di-toolkit==0.3.0 - - di-treetensor==0.4.1 - - diffusers==0.30.0 - - dill==0.3.8 - - diskcache==5.6.3 - - dm-control==1.0.22 - - dm-env==1.6 - - dm-tree==0.1.8 - - dnspython==2.6.1 - - docker-pycreds==0.4.0 - - docstring-parser==0.17.0 - - easydict==1.9 - - einops==0.8.1 - - email-validator==2.3.0 - - en-core-web-sm==3.8.0 - - enum-tools==0.12.0 - - etils==1.7.0 - - expecttest==0.2.1 - - farama-notifications==0.0.4 - - fastapi==0.119.1 - - fastapi-cli==0.0.13 - - fastapi-cloud-cli==0.3.1 - - fasteners==0.19 - - fastrlock==0.8.3 - - filelock==3.20.0 - - flask==2.0.3 - - fonttools==4.53.1 - - frozenlist==1.4.1 - - fsspec==2024.6.0 - - gguf==0.17.1 - - gitdb==4.0.11 - - gitpython==3.1.43 - - glfw==2.7.0 - - grpcio==1.75.1 - - gym==0.25.1 - - gym-notices==0.0.8 - - gymnasium==0.28.0 - - h11==0.16.0 - - h5py==3.11.0 - - hbutils==0.10.0 - - hf-xet==1.1.10 - - hickle==5.0.3 - - httpcore==1.0.9 - - httptools==0.7.1 - - httpx==0.28.1 - - huggingface-hub==0.35.3 - - hypothesis==6.103.0 - - imageio==2.35.1 - - imageio-ffmpeg==0.5.1 - - importlib-metadata==8.4.0 - - importlib-resources==6.4.4 - - iniconfig==2.3.0 - - interegular==0.3.3 - - ipykernel==6.29.5 - - ipywidgets==8.1.3 - - itsdangerous==2.2.0 - - jax-jumpy==1.0.0 - - jericho==3.3.0 - - jinja2==3.1.6 - - jiter==0.11.1 - - joblib==1.4.2 - - jsonschema==4.25.1 - - jupyter-client==8.6.2 - - jupyter-core==5.7.2 - - jupyterlab-widgets==3.0.11 - - kiwisolver==1.4.5 - - labmaze==1.0.6 - - langcodes==3.5.0 - - language-data==1.3.0 - - lark==1.2.2 - - lightning-utilities==0.11.6 - - lightzero==0.2.0 - - line-profiler==5.0.0 - - llguidance==0.7.30 - - llvmlite==0.44.0 - - lm-format-enforcer==0.11.3 - - lockfile==0.12.2 - - loguru==0.7.3 - - lxml==5.3.0 - - marisa-trie==1.3.1 - - markdown==3.9 - - markdown-it-py==3.0.0 - - matplotlib==3.9.2 - - mdurl==0.1.2 - - minigrid==2.2.1 - - mistral-common==1.8.5 - - mjrl==1.0.0 - - moviepy==1.0.3 - - mpire==2.10.2 - - msgpack==1.1.2 - - msgspec==0.19.0 - - mujoco==3.2.2 - - mujoco-py==2.1.2.14 - - multidict==6.0.5 - - multiprocess==0.70.16 - - murmurhash==1.0.13 - - nest-asyncio==1.6.0 - - networkx==3.3 - - ninja==1.13.0 - - nltk==3.9.2 - - numba==0.61.2 - - nvidia-cublas-cu12==12.8.4.1 - - nvidia-cuda-cupti-cu12==12.8.90 - - nvidia-cuda-nvrtc-cu12==12.8.93 - - nvidia-cuda-runtime-cu12==12.8.90 - - nvidia-cudnn-cu12==9.10.2.21 - - nvidia-cufft-cu12==11.3.3.83 - - nvidia-cufile-cu12==1.13.1.3 - - nvidia-curand-cu12==10.3.9.90 - - nvidia-cusolver-cu12==11.7.3.90 - - nvidia-cusparse-cu12==12.5.8.93 - - nvidia-cusparselt-cu12==0.7.1 - - nvidia-ml-py==13.580.82 - - nvidia-nccl-cu12==2.27.3 - - nvidia-nvjitlink-cu12==12.8.93 - - nvidia-nvtx-cu12==12.8.90 - - nvitop==1.5.3 - - openai==2.5.0 - - openai-harmony==0.0.4 - - opencv-python==4.10.0.84 - - opencv-python-headless==4.12.0.88 - - optree==0.11.0 - - orjson==3.10.7 - - outlines-core==0.2.11 - - pandas==2.3.3 - - partial-json-parser==0.2.1.1.post6 - - pastel==0.2.1 - - peft==0.17.1 - - pip==24.2 - - pluggy==1.6.0 - - poethepoet==0.10.0 - - pot==0.9.4 - - preshed==3.0.10 - - proglog==0.1.10 - - prometheus-client==0.23.1 - - prometheus-fastapi-instrumentator==7.1.0 - - protobuf==5.27.3 - - py-cpuinfo==9.0.0 - - pyarrow==21.0.0 - - pybase64==1.4.2 - - pybullet==3.2.6 - - pycountry==24.6.1 - - pydantic==2.12.3 - - pydantic-core==2.41.4 - - pydantic-extra-types==2.10.6 - - pygame==2.6.1 - - pympler==1.1 - - pynng==0.8.1 - - pyopengl==3.1.7 - - pyparsing==3.1.2 - - pytest==8.4.2 - - python-dateutil==2.9.0.post0 - - python-dotenv==1.1.1 - - python-etcd==0.4.5 - - python-graphviz==0.20.3 - - python-json-logger==4.0.0 - - python-multipart==0.0.20 - - pytimeparse==1.1.8 - - pytorch-lightning==2.4.0 - - pyzmq==26.2.0 - - ray==2.50.1 - - redis==6.4.0 - - regex==2024.7.24 - - responses==0.25.8 - - rich==13.7.1 - - rich-toolkit==0.15.1 - - rignore==0.7.1 - - safetensors==0.4.4 - - scikit-learn==1.5.1 - - scipy==1.14.1 - - seaborn==0.13.2 - - sentencepiece==0.2.1 - - sentry-sdk==2.42.0 - - setproctitle==1.3.3 - - setuptools==66.1.1 - - shellingham==1.5.4 - - shimmy==0.2.1 - - simple-parsing==0.1.7 - - smart-open==7.4.0 - - smmap==5.0.1 - - sniffio==1.3.1 - - sortedcontainers==2.4.0 - - soundfile==0.13.1 - - soxr==1.0.0 - - spacy==3.8.7 - - spacy-legacy==3.0.12 - - spacy-loggers==1.0.5 - - srsly==2.5.1 - - starlette==0.48.0 - - sympy==1.14.0 - - tabulate==0.9.0 - - tensorboard==2.20.0 - - tensorboard-data-server==0.7.2 - - tensorboardx==2.6.4 - - tensordict==0.5.0 - - termcolor==2.4.0 - - thinc==8.3.6 - - threadpoolctl==3.5.0 - - tiktoken==0.12.0 - - tokenizers==0.22.1 - - tomlkit==0.13.2 - - torch==2.8.0 - - torchaudio==2.8.0 - - torchcde==0.2.5 - - torchdiffeq==0.2.4 - - torchelastic==0.2.2 - - torchmetrics==1.4.1 - - torchsde==0.2.6 - - torchvision==0.23.0 - - tornado==6.4.1 - - trampoline==0.1.2 - - transformers==4.57.1 - - treevalue==1.4.12 - - triton==3.4.0 - - trueskill==0.4.5 - - typer==0.19.2 - - types-dataclasses==0.6.6 - - typing-extensions==4.15.0 - - typing-inspection==0.4.2 - - tzdata==2025.2 - - urlobject==3.0.0 - - uvicorn==0.38.0 - - uvloop==0.22.1 - - vllm==0.11.0 - - wandb==0.17.7 - - wasabi==1.1.3 - - watchfiles==1.1.1 - - weasel==0.4.1 - - websockets==15.0.1 - - werkzeug==2.0.3 - - widgetsnbextension==4.0.11 - - wrapt==2.0.0 - - xformers==0.0.32.post1 - - xgrammar==0.1.25 - - xxhash==3.6.0 - - yapf==0.29.0 - - yarl==1.9.4 - - yattag==1.16.1 - - zipp==3.20.0 -prefix: /opt/conda diff --git a/lzero/entry/train_unizero.py b/lzero/entry/train_unizero.py index d09f963b7..2553b28fc 100644 --- a/lzero/entry/train_unizero.py +++ b/lzero/entry/train_unizero.py @@ -168,11 +168,7 @@ def train_unizero( # Evaluate policy performance if learner.train_iter == 0 or evaluator.should_eval(learner.train_iter): logging.info(f"Training iteration {learner.train_iter}: Starting evaluation...") - stop, reward = evaluator.eval(learner.save_checkpoint, learner.train_iter, collector.envstep) - logging.info(f"Training iteration {learner.train_iter}: Evaluation completed, stop condition: {stop}, current reward: {reward}") - if stop: - logging.info("Stopping condition met, training ends!") - break + _ = evaluator.eval(learner.save_checkpoint, learner.train_iter, collector.envstep) # Collect new data new_data = collector.collect(train_iter=learner.train_iter, policy_kwargs=collect_kwargs) diff --git a/lzero/entry/train_unizero_segment.py b/lzero/entry/train_unizero_segment.py index 0559934c0..ef964c84d 100644 --- a/lzero/entry/train_unizero_segment.py +++ b/lzero/entry/train_unizero_segment.py @@ -20,6 +20,7 @@ from lzero.policy import visit_count_temperature from lzero.policy.random_policy import LightZeroRandomPolicy from lzero.worker import MuZeroEvaluator as Evaluator +from lzero.worker import MuZeroPerLevelEvaluator from lzero.worker import MuZeroSegmentCollector as Collector from .utils import random_collect, calculate_update_per_collect @@ -97,7 +98,8 @@ def train_unizero_segment( replay_buffer = GameBuffer(policy_config) collector = Collector(env=collector_env, policy=policy.collect_mode, tb_logger=tb_logger, exp_name=cfg.exp_name, policy_config=policy_config) - evaluator = Evaluator(eval_freq=cfg.policy.eval_freq, n_evaluator_episode=cfg.env.n_evaluator_episode, + EvaluatorCls = MuZeroPerLevelEvaluator if cfg.policy.get('eval_per_level', False) else Evaluator + evaluator = EvaluatorCls(eval_freq=cfg.policy.eval_freq, n_evaluator_episode=cfg.env.n_evaluator_episode, stop_value=cfg.env.stop_value, env=evaluator_env, policy=policy.eval_mode, tb_logger=tb_logger, exp_name=cfg.exp_name, policy_config=policy_config) @@ -115,6 +117,7 @@ def train_unizero_segment( # TODO: for visualize # stop, reward = evaluator.eval(learner.save_checkpoint, learner.train_iter, collector.envstep) + evaluator.eval(learner.save_checkpoint, learner.train_iter, collector.envstep) buffer_reanalyze_count = 0 train_epoch = 0 @@ -157,9 +160,7 @@ def train_unizero_segment( # if learner.train_iter == 0 or evaluator.should_eval(learner.train_iter): if learner.train_iter > 0 and evaluator.should_eval(learner.train_iter): - stop, reward = evaluator.eval(learner.save_checkpoint, learner.train_iter, collector.envstep) - if stop: - break + evaluator.eval(learner.save_checkpoint, learner.train_iter, collector.envstep) # Collect new data new_data = collector.collect(train_iter=learner.train_iter, policy_kwargs=collect_kwargs) diff --git a/lzero/entry/utils.py b/lzero/entry/utils.py index 99b22b852..0ec97a12c 100644 --- a/lzero/entry/utils.py +++ b/lzero/entry/utils.py @@ -528,9 +528,11 @@ def calculate_update_per_collect( collected_transitions_tensor ).item() updates = int(total_collected_transitions * cfg.policy.replay_ratio) + print(f"\ntotal_collected_transitions={total_collected_transitions}\tupdates={updates}\n") else: # In a single-process setup. updates = int(collected_transitions_num * cfg.policy.replay_ratio) + print(f"collected_transitions_num={collected_transitions_num}\tupdates={updates}") return max(1, updates) # Ensure at least one update. diff --git a/lzero/mcts/buffer/__init__.py b/lzero/mcts/buffer/__init__.py index d7ccb0678..541dd35c5 100644 --- a/lzero/mcts/buffer/__init__.py +++ b/lzero/mcts/buffer/__init__.py @@ -8,3 +8,4 @@ from .game_buffer_stochastic_muzero import StochasticMuZeroGameBuffer from .game_buffer_rezero_mz import ReZeroMZGameBuffer from .game_buffer_rezero_ez import ReZeroEZGameBuffer +from .game_buffer_priorzero import PriorZeroGameBufferOptimized diff --git a/lzero/mcts/buffer/game_buffer.py b/lzero/mcts/buffer/game_buffer.py index 253935652..d1a854988 100644 --- a/lzero/mcts/buffer/game_buffer.py +++ b/lzero/mcts/buffer/game_buffer.py @@ -145,14 +145,6 @@ def _sample_orig_data(self, batch_size: int, print_priority_logs: bool = False) game_segment = self.game_segment_buffer[game_segment_idx] game_segment_list.append(game_segment) - - # print(f'len(game_segment)=:len(game_segment.action_segment): {len(game_segment)}') - # print(f'len(game_segment.obs_segment): {game_segment.obs_segment.shape[0]}') - - # In the reanalysis phase, `pos_in_game_segment` should be a multiple of `num_unroll_steps`. - # Indices exceeding `game_segment_length` are padded with the next segment and are not updated - # in the current implementation. Therefore, we need to sample `pos_in_game_segment` within - # [0, game_segment_length - num_unroll_steps] to avoid padded data. if self._cfg.action_type == 'varied_action_space': # For some environments (e.g., Jericho), the action space size may be different. diff --git a/lzero/mcts/buffer/game_buffer_priorzero.py b/lzero/mcts/buffer/game_buffer_priorzero.py index c9dda2cf5..a408b0ebe 100644 --- a/lzero/mcts/buffer/game_buffer_priorzero.py +++ b/lzero/mcts/buffer/game_buffer_priorzero.py @@ -1,210 +1,73 @@ -# game_buffer_priorzero.py -""" -[PRIORZERO] Enhanced Game Buffer for PriorZero - -This module extends UniZeroGameBuffer to support LLM policy training (SFT + RFT). - -Key Features: -- Returns game_segments in sample() for LLM training data extraction -- Efficient indexing to avoid duplicating large observation data -- Robust handling of edge cases (partial batches, variable-length segments) -- Minimal memory overhead (only stores references, not copies) - -Author: PriorZero Team -Date: 2025-01-21 -""" - import numpy as np from typing import List, Any, Union, Tuple from lzero.mcts.buffer.game_buffer_unizero import UniZeroGameBuffer +from lzero.policy import to_detach_cpu_numpy, concat_output_value, inverse_scalar_transform +from lzero.mcts.utils import prepare_observation +import torch -class PriorZeroGameBuffer(UniZeroGameBuffer): - """ - [PRIORZERO-MODIFIED] - Enhanced GameBuffer that provides game_segments for LLM policy training. - - Modifications: - 1. sample() returns game_segments as 4th element - 2. Efficient implementation using existing game_segment_list from _make_batch - 3. No additional memory overhead (returns references, not copies) - """ +class PriorZeroGameBufferOptimized(UniZeroGameBuffer): def __init__(self, cfg): - """Initialize PriorZero Game Buffer.""" super().__init__(cfg) - - # [PRIORZERO-NEW] Cache for the last sampled game segments - # This avoids re-sampling when we need game segments - self._last_sampled_game_segments = None - self._last_sampled_batch_indices = None - - def sample( - self, - batch_size: int, - policy: Union["MuZeroPolicy", "EfficientZeroPolicy", "SampledEfficientZeroPolicy"] - ) -> List[Any]: + self.last_pos_in_transition = 0 + + def mark_latest_transitions_consumed(self) -> None: + self.last_pos_in_transition = self.get_num_of_transitions() + + def fetch_latest_batch(self, batch_size: int, policy, select_last: bool) -> List[Any]: """ - [PRIORZERO-MODIFIED] - Sample data and prepare current_batch, target_batch, AND game_segments. + Fetch latest batch for LLM training. Returns: - train_data: [current_batch, target_batch, game_segments] - - current_batch: [obs, action, target_action, mask, indices, weights, make_time, timestep] - - target_batch: [rewards, values, policies] - - game_segments: List of GameSegment objects used in this batch - - Note: - game_segments are returned for LLM training (SFT/RFT). - They contain: - - mcts_policy_segment: MCTS visit distributions (for SFT supervision) - - raw_obs_segment: Raw text observations (for LLM prompts) - - reward_segment: Environment rewards (for RFT) - - search_value_segment: MCTS search values (for analysis) + [raw_obs_list, history_obs_list, llm_prior_per_tok_list, + batch_target_values, batch_pred_values, cot_prefix_list, llm_action_list, action_list] + action_list: integer action indices for correct rollout log-prob lookup in VL training. """ policy._target_model.to(self._cfg.device) policy._target_model.eval() - # ====================================================================== - # [PRIORZERO-KEY] Sample data and extract game_segments - # ====================================================================== - # obtain the current_batch and prepare target context reward_value_context, policy_re_context, policy_non_re_context, current_batch = self._make_batch( - batch_size, self._cfg.reanalyze_ratio + batch_size, self._cfg.reanalyze_ratio, fetch_latest=True, select_last=select_last ) + if not current_batch: + return [[], [], [], [], [], [], [], []] - # [PRIORZERO-NEW] Extract game_segments from the sampling process - # These were already created in _make_batch, we just need to save them - game_segments = self._last_sampled_game_segments - - # Defensive check: ensure game_segments match batch_size - if game_segments is None or len(game_segments) != len(current_batch[4]): # current_batch[4] is batch_index_list - # Fallback: create empty list if something went wrong - import logging - logging.warning( - f"[PriorZeroBuffer] game_segments mismatch: " - f"expected {len(current_batch[4])}, got {len(game_segments) if game_segments else None}. " - f"Falling back to empty list (SFT/RFT will be skipped)." - ) - game_segments = [] - - # ====================================================================== - # Standard UniZero processing (unchanged) - # ====================================================================== - # current_batch = [obs_list, action_list, bootstrap_action_list, mask_list, batch_index_list, weights_list, make_time_list, timestep_list] + obs_list, action_list, bootstrap_action_list, mask_list, batch_index_list, weights_list, make_time_list, timestep_list, raw_obs_list, history_obs_list, llm_prior_per_tok_list, cot_prefix_list, llm_action_list = current_batch - # target reward, target value - batch_rewards, batch_target_values = self._compute_target_reward_value( - reward_value_context, policy._target_model, current_batch[2], current_batch[-1] # current_batch[2] is batch_target_action + # Standard processing + batch_rewards, batch_target_values, batch_pred_values = self._compute_target_reward_value_and_pred_value( + reward_value_context, policy._target_model, action_list, bootstrap_action_list, timestep_list ) - # target policy - batch_target_policies_re = self._compute_target_policy_reanalyzed( - policy_re_context, policy._target_model, current_batch[1], current_batch[-1] - ) # current_batch[1] is batch_action - batch_target_policies_non_re = self._compute_target_policy_non_reanalyzed( + batch_target_policies = self._compute_target_policy_non_reanalyzed( policy_non_re_context, self.action_space_size ) - # fusion of batch_target_policies_re and batch_target_policies_non_re to batch_target_policies - if 0 < self._cfg.reanalyze_ratio < 1: - batch_target_policies = np.concatenate([batch_target_policies_re, batch_target_policies_non_re]) - elif self._cfg.reanalyze_ratio == 1: - batch_target_policies = batch_target_policies_re - elif self._cfg.reanalyze_ratio == 0: - batch_target_policies = batch_target_policies_non_re - - target_batch = [batch_rewards, batch_target_values, batch_target_policies] - - # ====================================================================== - # [PRIORZERO-KEY] Return current_batch, target_batch, AND game_segments - # ====================================================================== - train_data = [current_batch, target_batch, game_segments] - return train_data - - def _sample_orig_data(self, batch_size: int) -> Tuple[Any]: - """ - [PRIORZERO-MODIFIED] - Override to cache game_segments during sampling. - - This avoids double sampling by caching the result for sample() to use. - """ - # Call parent implementation - result = super()._sample_orig_data(batch_size) - - # Cache the game_segment_list (first element of result tuple) - game_segment_list = result[0] - self._last_sampled_game_segments = game_segment_list - self._last_sampled_batch_indices = result[2] # batch_index_list + # CoT reuse optimization: return cot_prefix_list + # IMPORTANT: Validate return value before returning to ensure broadcast compatibility + # action_list included so VL training can index into action_logprobs by MCTS-selected action index + result = [raw_obs_list, history_obs_list, llm_prior_per_tok_list, batch_target_values, batch_pred_values, cot_prefix_list, llm_action_list, action_list] return result - - def _sample_orig_data_episode(self, batch_size: int) -> Tuple[Any]: - """ - [PRIORZERO-MODIFIED] - Override to cache game_segments during episode sampling. - - This avoids double sampling by caching the result for sample() to use. - """ - # Call parent implementation - result = super()._sample_orig_data_episode(batch_size) - - # Cache the game_segment_list (first element of result tuple) - game_segment_list = result[0] - self._last_sampled_game_segments = game_segment_list - self._last_sampled_batch_indices = result[2] # batch_index_list - - return result - - def clear(self): - """ - [PRIORZERO-MODIFIED] - Clear buffer and cached game segments. - """ - super().clear() - self._last_sampled_game_segments = None - self._last_sampled_batch_indices = None - - -# ============================================================================== -# Optimized Alternative (Avoids Double Sampling) -# ============================================================================== - -class PriorZeroGameBufferOptimized(UniZeroGameBuffer): - """ - [PRIORZERO-OPTIMIZED] - More efficient version that avoids double sampling by modifying _make_batch minimally. - - This version uses a monkey-patch approach to intercept orig_data during parent's _make_batch call. - """ - - def __init__(self, cfg): - super().__init__(cfg) - self._cached_game_segments = None - + def sample(self, batch_size: int, policy) -> List[Any]: """Sample data with game_segments (optimized version).""" policy._target_model.to(self._cfg.device) policy._target_model.eval() - # Reset cache - self._cached_game_segments = None - - # Call parent's _make_batch (which will trigger our hook) reward_value_context, policy_re_context, policy_non_re_context, current_batch = self._make_batch( batch_size, self._cfg.reanalyze_ratio ) - # Get cached game segments (set by our overridden _make_batch) - game_segments = self._cached_game_segments or [] - + obs_list, action_list, bootstrap_action_list, mask_list, batch_index_list, weights_list, make_time_list, timestep_list, raw_obs_list, history_obs_list, llm_prior_per_tok_list, cot_prefix_list, llm_action_list = current_batch # Standard processing batch_rewards, batch_target_values = self._compute_target_reward_value( - reward_value_context, policy._target_model, current_batch[2], current_batch[-1] + reward_value_context, policy._target_model, current_batch[2], timestep_list ) batch_target_policies_re = self._compute_target_policy_reanalyzed( - policy_re_context, policy._target_model, current_batch[1], current_batch[-1] + policy_re_context, policy._target_model, current_batch[1], timestep_list ) batch_target_policies_non_re = self._compute_target_policy_non_reanalyzed( policy_non_re_context, self.action_space_size @@ -219,30 +82,33 @@ def sample(self, batch_size: int, policy) -> List[Any]: target_batch = [batch_rewards, batch_target_values, batch_target_policies] - return [current_batch, target_batch, game_segments] + return [current_batch, target_batch] - def _make_batch(self, batch_size: int, reanalyze_ratio: float) -> Tuple[Any]: - """ - [PRIORZERO-OPTIMIZED] - Minimally modified to cache game_segment_list during sampling. + def _make_batch(self, batch_size: int, reanalyze_ratio: float, fetch_latest: bool = False, select_last: bool = False) -> Tuple[Any]: - This is a full override of parent's _make_batch to avoid double sampling. - Code is mostly copied from parent, with one key addition: caching game_segments. - """ # Sample original data - if self.sample_type == 'transition': - orig_data = self._sample_orig_data(batch_size) - elif self.sample_type == 'episode': - orig_data = self._sample_orig_data_episode(batch_size) + if not fetch_latest: + if self.sample_type == 'transition': + orig_data = self._sample_orig_data(batch_size) + elif self.sample_type == 'episode': + orig_data = self._sample_orig_data_episode(batch_size) + else: + if self.sample_type == 'transition': + orig_data = self._fetch_latest_orig_data(batch_size, select_last=select_last) + elif self.sample_type == 'episode': + raise ValueError("fetch_latest with episode sampling not supported.") game_segment_list, pos_in_game_segment_list, batch_index_list, weights_list, make_time_list = orig_data - - # [PRIORZERO-KEY] Cache game_segments for sample() to use - self._cached_game_segments = game_segment_list - + if not pos_in_game_segment_list: + return [], [], [], [] + # Rest of the code is identical to parent's _make_batch batch_size = len(batch_index_list) obs_list, action_list, mask_list = [], [], [] + raw_obs_list, history_obs_list = [], [] + llm_prior_per_tok_list = [] + cot_prefix_list = [] # CoT reuse optimization + llm_action_list = [] timestep_list = [] bootstrap_action_list = [] @@ -272,6 +138,22 @@ def _make_batch(self, batch_size: int, reanalyze_ratio: float) -> Tuple[Any]: pos_in_game_segment_list[i], num_unroll_steps=self._cfg.num_unroll_steps, padding=True ) ) + raw_obs_list.append(game_segment_list[i].get_unroll_raw_obs( + pos_in_game_segment_list[i], num_unroll_steps=self._cfg.num_unroll_steps, padding=True + )) + history_obs_list.append(game_segment_list[i].get_unroll_histroy_obs( + pos_in_game_segment_list[i], num_unroll_steps=self._cfg.num_unroll_steps, padding=True + )) + llm_prior_per_tok_list.append(game_segment_list[i].get_unroll_llm_prior_per_tok( + pos_in_game_segment_list[i], num_unroll_steps=self._cfg.num_unroll_steps, padding=True + )) + cot_prefix_list.append(game_segment_list[i].get_unroll_cot_prefix( + pos_in_game_segment_list[i], num_unroll_steps=self._cfg.num_unroll_steps, padding=True + )) + llm_action_list.append(game_segment_list[i].get_unroll_llm_action( + pos_in_game_segment_list[i], num_unroll_steps=self._cfg.num_unroll_steps, padding=True + )) + action_list.append(actions_tmp) mask_list.append(mask_tmp) timestep_list.append(timestep_tmp) @@ -291,12 +173,54 @@ def _make_batch(self, batch_size: int, reanalyze_ratio: float) -> Tuple[Any]: current_batch = [obs_list, action_list, bootstrap_action_list, mask_list, batch_index_list, weights_list, make_time_list, timestep_list] for i in range(len(current_batch)): current_batch[i] = np.asarray(current_batch[i]) + # 检查 vllm和policy_model的输入上下文是否一致 (only for non-padded positions) + assert len(raw_obs_list) == len(history_obs_list) == len(llm_prior_per_tok_list) == len(cot_prefix_list) == len(llm_action_list) + B, T = len(raw_obs_list), len(raw_obs_list[0]) + # Only run dict-based consistency checks for LLM text path. + # In VL (image) mode, llm_prior_per_tok entries are numpy arrays (or None), not dicts. + # Additional 'prefix_cot' key check prevents false positives if VL path ever returns dicts. + _is_llm_text_mode = ( + B > 0 and T > 1 + and llm_prior_per_tok_list[0][1] is not None + and isinstance(llm_prior_per_tok_list[0][1], dict) + and 'prefix_cot' in llm_prior_per_tok_list[0][1] + ) + if _is_llm_text_mode: + for b in range(B): + for t in range(T - 1): + # Skip padded positions: mask[t] == 0 means the action at step t is padding, + # so llm_prior_per_tok at t+1 is also padding and the alignment invariant doesn't hold. + if mask_list[b][t] == 0.: + continue + current_obs = raw_obs_list[b][t] + current_hist = history_obs_list[b][t] + + old_prefix_cot = llm_prior_per_tok_list[b][t+1]['prefix_cot'] + old_current_obs = llm_prior_per_tok_list[b][t+1]['current_obs'] + old_history = llm_prior_per_tok_list[b][t+1]['history'] + old_logprob = llm_prior_per_tok_list[b][t+1]['rollout_action_logprob'] + cot_prefix = cot_prefix_list[b][t+1] + llm_action = llm_action_list[b][t+1] + + assert llm_action in old_logprob + assert old_current_obs == current_obs and old_history == current_hist and old_prefix_cot == cot_prefix + + current_batch.append(raw_obs_list) + current_batch.append(history_obs_list) + current_batch.append(llm_prior_per_tok_list) + current_batch.append(cot_prefix_list) # CoT reuse optimization + current_batch.append(llm_action_list) total_transitions = self.get_num_of_transitions() - reward_value_context = self._prepare_reward_value_context( - batch_index_list, game_segment_list, pos_in_game_segment_list, total_transitions - ) + if not fetch_latest: + reward_value_context = self._prepare_reward_value_context( + batch_index_list, game_segment_list, pos_in_game_segment_list, total_transitions + ) + else: + reward_value_context = self._prepare_reward_value_context_and_pred_values( + batch_index_list, game_segment_list, pos_in_game_segment_list, total_transitions + ) reanalyze_num = max(int(batch_size * reanalyze_ratio), 1) if reanalyze_ratio > 0 else 0 self.reanalyze_num = reanalyze_num @@ -319,71 +243,300 @@ def _make_batch(self, batch_size: int, reanalyze_ratio: float) -> Tuple[Any]: return reward_value_context, policy_re_context, policy_non_re_context, current_batch - -# ============================================================================== -# Factory Function -# ============================================================================== - -def create_priorzero_buffer(cfg, optimized: bool = True): - """ - Factory function to create PriorZero game buffer. - - Args: - cfg: Configuration dict - optimized: If True, use optimized version (recommended) - - Returns: - buffer: PriorZero game buffer instance - """ - if optimized: - return PriorZeroGameBufferOptimized(cfg) - else: - return PriorZeroGameBuffer(cfg) - - -if __name__ == "__main__": - print("="*80) - print("PriorZero Game Buffer - Unit Tests") - print("="*80) - - # Create mock config - class MockConfig: - def __init__(self): - self.device = 'cpu' - self.env_type = 'not_board_games' - self.game_segment_length = 200 - self.num_unroll_steps = 5 - self.td_steps = 5 - self.batch_size = 32 - self.use_priority = False - self.reanalyze_ratio = 0.0 - self.sample_type = 'transition' - self.replay_buffer_size = 10000 - self.model = type('obj', (object,), { - 'model_type': 'mlp', - 'action_space_size': 10, - 'observation_shape': 128, - })() - - cfg = MockConfig() - - # Test both versions - for name, buffer_class in [ - ("Standard", PriorZeroGameBuffer), - ("Optimized", PriorZeroGameBufferOptimized) - ]: - print(f"\nTesting {name} Buffer:") - print("-" * 40) - - buffer = buffer_class(cfg) - print(f"✓ Buffer created: {type(buffer).__name__}") - print(f" - sample_type: {buffer.sample_type}") - print(f" - action_space_size: {buffer.action_space_size}") - - # Note: Full testing would require mock GameSegments and Policy - # For now, just verify instantiation - print(f"✓ {name} buffer initialized successfully") - - print("\n" + "="*80) - print("✓ All tests passed!") - print("="*80) + def _clear(self): + self.game_pos_priorities = [] + self.game_segment_buffer = [] + self.game_segment_game_pos_look_up = [] + + + def _fetch_latest_orig_data(self, batch_size: int, select_last: bool = False) -> Tuple: + """ + Overview: + Sample original data which includes: + - game_segment_list: A list of game segments. + - pos_in_game_segment_list: Transition index in the game (relative index). + - batch_index_list: The index of the start transition of the sampled mini-batch in the replay buffer. + - weights_list: The weight concerning the priority. + - make_time: The time the batch is made (for correctly updating the replay buffer when data is deleted). + Arguments: + - batch_size (:obj:`int`): The size of the batch. + - print_priority_logs (:obj:`bool`): Whether to print logs related to priority statistics, defaults to False. + """ + assert self._beta > 0, "Beta should be greater than 0" + num_of_transitions = self.get_num_of_transitions() + + probs = self.game_pos_priorities ** self._alpha + 1e-6 + probs /= probs.sum() + + # 主要改动: 由sample改成了确定的取最后batch_size个样本 + if select_last: + latest_new_indices = list(range(self.last_pos_in_transition, num_of_transitions)) + if batch_size == -1: + candidate_batch_index_list = latest_new_indices + else: + candidate_batch_index_list = latest_new_indices[-batch_size:] + else: + latest_new_indices = list(range(num_of_transitions)) + candidate_batch_index_list = np.random.choice(num_of_transitions, size=batch_size, replace=False, p=probs) + game_segment_list = [] + pos_in_game_segment_list = [] + batch_index_list = [] + + for idx in candidate_batch_index_list: + game_segment_idx, pos_in_game_segment = self.game_segment_game_pos_look_up[idx] + game_segment_idx -= self.base_idx # Adjust index based on base index + game_segment = self.game_segment_buffer[game_segment_idx] + + assert len(game_segment.obs_segment) == len(game_segment.raw_obs_segment) == len(game_segment.cot_prefix_segment) + segment_len = len(game_segment.action_segment) + if self._cfg.action_type == 'varied_action_space': + within_obs_window = pos_in_game_segment + self._cfg.num_unroll_steps + self._cfg.model.frame_stack_num <= len(game_segment.obs_segment) + within_td_window = pos_in_game_segment < self._cfg.game_segment_length - self._cfg.num_unroll_steps + valid_next_action = pos_in_game_segment < segment_len - 1 + is_valid_latest_transition = within_obs_window and within_td_window and valid_next_action + else: + within_obs_window = pos_in_game_segment + self._cfg.num_unroll_steps + self._cfg.model.frame_stack_num <= len(game_segment.obs_segment) + within_segment_window = pos_in_game_segment < self._cfg.game_segment_length + valid_next_action = pos_in_game_segment < segment_len - 1 + is_valid_latest_transition = within_obs_window and within_segment_window and valid_next_action + + if not is_valid_latest_transition: + continue + + game_segment_list.append(game_segment) + pos_in_game_segment_list.append(pos_in_game_segment) + batch_index_list.append(idx) + + import random + n = min(256, len(game_segment_list)) + print(f"new transition={len(latest_new_indices)} | valid_pos_in_gamesemt={len(game_segment_list)} | final_pos_in_gamesemt={n}") + indices = random.sample(range(len(game_segment_list)), n) + game_segment_list = [game_segment_list[i] for i in indices] + pos_in_game_segment_list = [pos_in_game_segment_list[i] for i in indices] + batch_index_list = [batch_index_list[i] for i in indices] + # make_time = [time.time() for _ in range(len(batch_index_list))] + + # Set the make_time for each sample (set to 0 for now, but can be the actual time if needed). + make_time = [0. for _ in range(len(batch_index_list))] + + orig_data = (game_segment_list, pos_in_game_segment_list, batch_index_list, None, make_time) + + return orig_data + + # 从原来的_prepare_reward_value_context函数修改得到 + def _prepare_reward_value_context_and_pred_values( + self, batch_index_list: List[str], game_segment_list: List[Any], pos_in_game_segment_list: List[Any], + total_transitions: int + ) -> List[Any]: + """ + Overview: + prepare the context of rewards and values for calculating TD value target in reanalyzing part. + Arguments: + - batch_index_list (:obj:`list`): the index of start transition of sampled minibatch in replay buffer + - game_segment_list (:obj:`list`): list of game segments + - pos_in_game_segment_list (:obj:`list`): list of transition index in game_segment + - total_transitions (:obj:`int`): number of collected transitions + Returns: + - reward_value_context (:obj:`list`): value_obs_list, value_mask, pos_in_game_segment_list, rewards_list, game_segment_lens, + td_steps_list, action_mask_segment, to_play_segment + """ + zero_obs = game_segment_list[0].zero_obs() + + pred_obs_list = [] + pred_mask = [] + + value_obs_list = [] + # the value is valid or not (out of game_segment) + value_mask = [] + rewards_list = [] + game_segment_lens = [] + # for board games + action_mask_segment, to_play_segment = [], [] + + root_values = [] + + td_steps_list = [] + for game_segment, state_index in zip(game_segment_list, pos_in_game_segment_list): + game_segment_len = len(game_segment) + game_segment_lens.append(game_segment_len) + # original buffer td-steps + td_steps = np.clip(self._cfg.td_steps, 1, max(1, game_segment_len - state_index)).astype(np.int32) + + # prepare the corresponding observations for bootstrapped values o_{t+k} + # o[t+ td_steps, t + td_steps + stack frames + num_unroll_steps] + # t=2+3 -> o[2+3, 2+3+4+5] -> o[5, 14] + game_obs_pred = game_segment.get_unroll_obs(state_index, self._cfg.num_unroll_steps) + game_obs = game_segment.get_unroll_obs(state_index + td_steps, self._cfg.num_unroll_steps) + + rewards_list.append(game_segment.reward_segment) + + # for board games + action_mask_segment.append(game_segment.action_mask_segment) + to_play_segment.append(game_segment.to_play_segment) + + truncation_length = game_segment_len + + for current_index in range(state_index, state_index + self._cfg.num_unroll_steps + 1): + # get the bootstrapped target obs + td_steps_list.append(td_steps) + # index of bootstrapped obs o_{t+td_steps} + bootstrap_index = current_index + td_steps + + beg_index = current_index - state_index + end_index = beg_index + self._cfg.model.frame_stack_num + + if bootstrap_index < truncation_length: + value_mask.append(1) + # the stacked obs in time t + obs = game_obs[beg_index:end_index] + else: + value_mask.append(0) + obs = zero_obs + + if current_index < truncation_length: + pred_mask.append(1) + obs_pred = game_obs_pred[beg_index:end_index] + else: + pred_mask.append(0) + obs_pred = zero_obs + + value_obs_list.append(obs) + pred_obs_list.append(obs_pred) + + reward_value_context = [ + value_obs_list, value_mask, pos_in_game_segment_list, rewards_list, root_values, game_segment_lens, td_steps_list, + action_mask_segment, to_play_segment, pred_obs_list, pred_mask + ] + return reward_value_context + + # 从原来的_compute_target_reward_value函数修改得到 + def _compute_target_reward_value_and_pred_value(self, reward_value_context: List[Any], model: Any, batch_action_pred, batch_action, batch_timestep) -> Tuple[Any, Any]: + """ + Overview: + prepare reward and value targets from the context of rewards and values. + Arguments: + - reward_value_context (:obj:'list'): the reward value context + - model (:obj:'torch.tensor'):model of the target model + Returns: + - batch_value_prefixs (:obj:'np.ndarray): batch of value prefix + - batch_target_values (:obj:'np.ndarray): batch of value estimation + """ + value_obs_list, value_mask, pos_in_game_segment_list, rewards_list, root_values, game_segment_lens, td_steps_list, action_mask_segment, \ + to_play_segment, pred_obs_list, pred_mask = reward_value_context # noqa + # transition_batch_size = game_segment_batch_size * (num_unroll_steps+1) + transition_batch_size = len(value_obs_list) + + batch_target_values, batch_rewards, batch_pred_values = [], [], [] + with torch.no_grad(): + value_obs_list = prepare_observation(value_obs_list, self._cfg.model.model_type) + pred_obs_list = prepare_observation(pred_obs_list, self._cfg.model.model_type) + + network_output = [] + network_output_pred = [] + + batch_obs = torch.from_numpy(value_obs_list).to(self._cfg.device).float() + batch_obs_pred = torch.from_numpy(pred_obs_list).to(self._cfg.device).float() + + # =============== NOTE: The key difference with MuZero ================= + # calculate the bootstrapped value and target value + # NOTE: batch_obs(value_obs_list) is at t+td_steps, batch_action is at timestep t+td_steps + if self.task_id is not None: + # m_output = model.initial_inference(batch_obs, batch_action, start_pos=batch_timestep, task_id=self.task_id) + m_output = model.initial_inference(batch_obs, batch_action, task_id=self.task_id) + m_output_pred = model.initial_inference(batch_obs_pred, batch_action_pred, task_id=self.task_id) + + else: + m_output = model.initial_inference(batch_obs, batch_action, start_pos=batch_timestep) + m_output_pred = model.initial_inference(batch_obs_pred, batch_action_pred, start_pos=batch_timestep) + + # ====================================================================== + + # if not in training, obtain the scalars of the value/reward + [m_output.latent_state, m_output.value, m_output.policy_logits] = to_detach_cpu_numpy( + [ + m_output.latent_state, + inverse_scalar_transform(m_output.value, self.value_support), + m_output.policy_logits + ] + ) + [m_output_pred.latent_state, m_output_pred.value, m_output_pred.policy_logits] = to_detach_cpu_numpy( + [ + m_output_pred.latent_state, + inverse_scalar_transform(m_output_pred.value, self.value_support), + m_output_pred.policy_logits + ] + ) + + network_output.append(m_output) + network_output_pred.append(m_output_pred) + + if self._cfg.use_root_value: + value_numpy = np.array(root_values) + raise ValueError("error!!!") + else: + # use the predicted values + value_numpy = concat_output_value(network_output) + pred_numpy = concat_output_value(network_output_pred) + + # 不考虑 board_games的情况 + value_numpy = value_numpy.reshape(-1) * ( + np.array([self._cfg.discount_factor for _ in range(transition_batch_size)]) ** td_steps_list + ) + pred_numpy = pred_numpy.reshape(-1) + + value_numpy= value_numpy * np.array(value_mask) + value_list = value_numpy.tolist() + + pred_numpy = pred_numpy * np.array(pred_mask) + pred_list = pred_numpy.tolist() + + + horizon_id, value_index = 0, 0 + + for game_segment_len_non_re, reward_list, state_index, to_play_list in zip(game_segment_lens, rewards_list, + pos_in_game_segment_list, + to_play_segment): + target_values = [] + target_rewards = [] + pred_values = [] + base_index = state_index + + # =========== NOTE =============== + # if game_segment_len_non_re < self._cfg.game_segment_length: + # # The last segment of one episode, the target value of excess part should be 0 + # truncation_length = game_segment_len_non_re + # else: + # # game_segment_len is game_segment.action_segment.shape[0] + # # action_segment.shape[0] = reward_segment.shape[0] or action_segment.shape[0] = reward_segment.shape[0] + 1 + # truncation_length = game_segment_len_non_re + # assert reward_list.shape[0] + 1 == game_segment_len_non_re or reward_list.shape[0] == game_segment_len_non_re + + truncation_length = game_segment_len_non_re + + for current_index in range(state_index, state_index + self._cfg.num_unroll_steps + 1): + bootstrap_index = current_index + td_steps_list[value_index] + for i, reward in enumerate(reward_list[current_index:bootstrap_index]): + # 不考虑 board_games的情况 + value_list[value_index] += reward * self._cfg.discount_factor ** i + horizon_id += 1 + + # TODO: check the boundary condition + target_values.append(value_list[value_index]) + pred_values.append(pred_list[value_index]) + + if current_index < len(reward_list): + target_rewards.append(reward_list[current_index]) + else: + target_rewards.append(np.array(0.)) + + value_index += 1 + + batch_rewards.append(target_rewards) + batch_target_values.append(target_values) + batch_pred_values.append(pred_values) + + batch_rewards = np.asarray(batch_rewards) + batch_target_values = np.asarray(batch_target_values) + batch_pred_values = np.asarray(batch_pred_values) + + return batch_rewards, batch_target_values, batch_pred_values \ No newline at end of file diff --git a/lzero/mcts/buffer/game_buffer_unizero.py b/lzero/mcts/buffer/game_buffer_unizero.py index 3bb9bf2ca..03180a24b 100644 --- a/lzero/mcts/buffer/game_buffer_unizero.py +++ b/lzero/mcts/buffer/game_buffer_unizero.py @@ -540,8 +540,7 @@ def _compute_target_policy_reanalyzed(self, policy_re_context: List[Any], model: return batch_target_policies_re - def _compute_target_reward_value(self, reward_value_context: List[Any], model: Any, batch_action, batch_timestep) -> Tuple[ - Any, Any]: + def _compute_target_reward_value(self, reward_value_context: List[Any], model: Any, batch_action, batch_timestep) -> Tuple[Any, Any]: """ Overview: prepare reward and value targets from the context of rewards and values. @@ -609,6 +608,7 @@ def _compute_target_reward_value(self, reward_value_context: List[Any], model: A value_numpy= value_numpy * np.array(value_mask) value_list = value_numpy.tolist() + horizon_id, value_index = 0, 0 for game_segment_len_non_re, reward_list, state_index, to_play_list in zip(game_segment_lens, rewards_list, diff --git a/lzero/mcts/buffer/game_segment.py b/lzero/mcts/buffer/game_segment.py index 2c45b328b..8dfa54dc4 100644 --- a/lzero/mcts/buffer/game_segment.py +++ b/lzero/mcts/buffer/game_segment.py @@ -150,6 +150,7 @@ def append( to_play: int = -1, timestep: int = 0, chance: int = 0, + **kwargs, ) -> None: """ Overview: diff --git a/lzero/mcts/ctree/ctree_efficientzero/lib/cnode.cpp b/lzero/mcts/ctree/ctree_efficientzero/lib/cnode.cpp index 8a94a7ca9..969b7a7d5 100644 --- a/lzero/mcts/ctree/ctree_efficientzero/lib/cnode.cpp +++ b/lzero/mcts/ctree/ctree_efficientzero/lib/cnode.cpp @@ -901,7 +901,6 @@ namespace tree get_time_and_set_rand_seed(); int last_action = -1; - float parent_q = 0.0; results.search_lens = std::vector(); int players = 0; @@ -917,6 +916,7 @@ namespace tree for (int i = 0; i < results.num; ++i) { + float parent_q = 0.0; CNode *node = &(roots->roots[i]); int is_root = 1; int search_len = 0; @@ -982,7 +982,6 @@ namespace tree get_time_and_set_rand_seed(); int last_action = -1; - float parent_q = 0.0; results.search_lens = std::vector(); int players = 0; @@ -998,6 +997,7 @@ namespace tree for (int i = 0; i < results.num; ++i) { + float parent_q = 0.0; CNode *node = &(roots->roots[i]); int is_root = 1; int search_len = 0; diff --git a/lzero/mcts/ctree/ctree_gumbel_muzero/lib/cnode.cpp b/lzero/mcts/ctree/ctree_gumbel_muzero/lib/cnode.cpp index 1adf1c1d2..5b8a4ef07 100644 --- a/lzero/mcts/ctree/ctree_gumbel_muzero/lib/cnode.cpp +++ b/lzero/mcts/ctree/ctree_gumbel_muzero/lib/cnode.cpp @@ -850,7 +850,6 @@ namespace tree{ srand(t1.tv_usec); int last_action = -1; - float parent_q = 0.0; results.search_lens = std::vector(); int players = 0; diff --git a/lzero/mcts/ctree/ctree_muzero/lib/cnode.cpp b/lzero/mcts/ctree/ctree_muzero/lib/cnode.cpp index 24eb3605c..63ae27b15 100644 --- a/lzero/mcts/ctree/ctree_muzero/lib/cnode.cpp +++ b/lzero/mcts/ctree/ctree_muzero/lib/cnode.cpp @@ -770,7 +770,6 @@ namespace tree get_time_and_set_rand_seed(); int last_action = -1; - float parent_q = 0.0; results.search_lens = std::vector(); int players = 0; @@ -782,6 +781,7 @@ namespace tree for (int i = 0; i < results.num; ++i) { + float parent_q = 0.0; CNode *node = &(roots->roots[i]); int is_root = 1; int search_len = 0; @@ -845,7 +845,6 @@ namespace tree get_time_and_set_rand_seed(); int last_action = -1; - float parent_q = 0.0; results.search_lens = std::vector(); int players = 0; @@ -857,6 +856,7 @@ namespace tree for (int i = 0; i < results.num; ++i) { + float parent_q = 0.0; CNode *node = &(roots->roots[i]); int is_root = 1; int search_len = 0; diff --git a/lzero/mcts/ctree/ctree_sampled_efficientzero/lib/cnode.cpp b/lzero/mcts/ctree/ctree_sampled_efficientzero/lib/cnode.cpp index 6bc4ea2e8..563b4f4c9 100644 --- a/lzero/mcts/ctree/ctree_sampled_efficientzero/lib/cnode.cpp +++ b/lzero/mcts/ctree/ctree_sampled_efficientzero/lib/cnode.cpp @@ -1132,7 +1132,6 @@ namespace tree } // CAction last_action = CAction(null_value, 1); std::vector last_action; - float parent_q = 0.0; results.search_lens = std::vector(); int players = 0; @@ -1144,6 +1143,7 @@ namespace tree for (int i = 0; i < results.num; ++i) { + float parent_q = 0.0; CNode *node = &(roots->roots[i]); int is_root = 1; int search_len = 0; diff --git a/lzero/mcts/ctree/ctree_sampled_muzero/lib/cnode.cpp b/lzero/mcts/ctree/ctree_sampled_muzero/lib/cnode.cpp index 83f50e2da..6ba2f96e4 100644 --- a/lzero/mcts/ctree/ctree_sampled_muzero/lib/cnode.cpp +++ b/lzero/mcts/ctree/ctree_sampled_muzero/lib/cnode.cpp @@ -1124,7 +1124,6 @@ namespace tree null_value.push_back(i + 0.1); } std::vector last_action; - float parent_q = 0.0; results.search_lens = std::vector(); int players = 0; @@ -1136,6 +1135,7 @@ namespace tree for (int i = 0; i < results.num; ++i) { + float parent_q = 0.0; CNode *node = &(roots->roots[i]); int is_root = 1; int search_len = 0; diff --git a/lzero/mcts/ptree/ptree_ez.py b/lzero/mcts/ptree/ptree_ez.py index 0a3058b26..32d3b72cb 100644 --- a/lzero/mcts/ptree/ptree_ez.py +++ b/lzero/mcts/ptree/ptree_ez.py @@ -474,7 +474,6 @@ def batch_traverse( - virtual_to_play (:obj:`Union[list, int]`): The to_play list used in self_play collecting and trainin gin board games, `virtual` is to emphasize that actions are performed on an imaginary hidden state. """ - parent_q = 0.0 results.search_lens = [None for _ in range(results.num)] results.last_actions = [None for _ in range(results.num)] results.nodes = [None for _ in range(results.num)] @@ -494,6 +493,7 @@ def batch_traverse( players = 1 for i in range(results.num): + parent_q = 0.0 node = roots.roots[i] is_root = 1 search_len = 0 diff --git a/lzero/mcts/ptree/ptree_mz.py b/lzero/mcts/ptree/ptree_mz.py index 794cd002a..e8f143560 100644 --- a/lzero/mcts/ptree/ptree_mz.py +++ b/lzero/mcts/ptree/ptree_mz.py @@ -446,7 +446,6 @@ def batch_traverse( - virtual_to_play (:obj:`list`): The to_play list used in self_play collecting and trainin gin board games, `virtual` is to emphasize that actions are performed on an imaginary hidden state. """ - parent_q = 0.0 results.search_lens = [None for _ in range(results.num)] results.last_actions = [None for _ in range(results.num)] @@ -460,6 +459,7 @@ def batch_traverse( results.search_paths = {i: [] for i in range(results.num)} for i in range(results.num): + parent_q = 0.0 node = roots.roots[i] is_root = 1 search_len = 0 diff --git a/lzero/mcts/ptree/ptree_sez.py b/lzero/mcts/ptree/ptree_sez.py index 6262891dc..63d1dde2e 100644 --- a/lzero/mcts/ptree/ptree_sez.py +++ b/lzero/mcts/ptree/ptree_sez.py @@ -666,7 +666,6 @@ def batch_traverse( - virtual_to_play (:obj:`list`): The to_play list used in self_play collecting and trainin gin board games, `virtual` is to emphasize that actions are performed on an imaginary hidden state. """ - parent_q = 0.0 results.search_lens = [None for _ in range(results.num)] results.last_actions = [None for _ in range(results.num)] @@ -680,6 +679,7 @@ def batch_traverse( results.search_paths = {i: [] for i in range(results.num)} for i in range(results.num): + parent_q = 0.0 node = roots.roots[i] is_root = 1 search_len = 0 diff --git a/lzero/mcts/ptree/ptree_stochastic_mz.py b/lzero/mcts/ptree/ptree_stochastic_mz.py index e407cfcde..6c9de5567 100644 --- a/lzero/mcts/ptree/ptree_stochastic_mz.py +++ b/lzero/mcts/ptree/ptree_stochastic_mz.py @@ -486,7 +486,6 @@ def batch_traverse( - virtual_to_play (:obj:`list`): The to_play list used in self_play collecting and trainin gin board games, `virtual` is to emphasize that actions are performed on an imaginary hidden state. """ - parent_q = 0.0 results.search_lens = [None for i in range(results.num)] results.last_actions = [None for i in range(results.num)] @@ -500,6 +499,7 @@ def batch_traverse( results.search_paths = {i: [] for i in range(results.num)} for i in range(results.num): + parent_q = 0.0 node = roots.roots[i] is_root = 1 search_len = 0 diff --git a/lzero/model/common.py b/lzero/model/common.py index 8c7bcdef2..f326724db 100644 --- a/lzero/model/common.py +++ b/lzero/model/common.py @@ -504,6 +504,9 @@ def __init__(self, torch.distributed.barrier() if get_rank() != 0: self.pretrained_model = AutoModel.from_pretrained(model_path) + + for p in self.pretrained_model.parameters(): + p.requires_grad = False self.embedding_size = embedding_size self.embed_proj_head = nn.Linear(self.pretrained_model.config.hidden_size, self.embedding_size) @@ -526,8 +529,22 @@ def forward(self, x: torch.Tensor, no_grad: bool = True) -> torch.Tensor: Returns: - (:obj:`torch.Tensor`): The final language embedding of shape (B, embedding_size). """ + # Ensure the input has a batch dimension for BERT. + if x.dim() == 1: + x = x.unsqueeze(0) # Ensure the input tensor is of type long. x = x.long() + + # Guard: BERT requires seq_len > 0. Return a zero embedding when the + # input is degenerate (empty batch or zero-length sequence). + if x.numel() == 0 or (x.dim() >= 2 and x.shape[1] == 0): + import logging + logging.getLogger(__name__).warning( + f"[HFLanguageRepresentationNetwork] Empty input detected: x.shape={x.shape}. " + "Returning zero embeddings." + ) + batch = x.shape[0] if x.dim() >= 2 else 1 + return torch.zeros(batch, self.embed_proj_head.out_features, device=x.device) # Construct the attention mask to exclude padding tokens. attention_mask = (x != self.tokenizer.pad_token_id).long() diff --git a/lzero/model/unizero_world_models/hf_transformer.py b/lzero/model/unizero_world_models/hf_transformer.py new file mode 100644 index 000000000..c32c9e440 --- /dev/null +++ b/lzero/model/unizero_world_models/hf_transformer.py @@ -0,0 +1,86 @@ +from typing import Optional + +import torch +from transformers import Qwen2ForCausalLM +from transformers.cache_utils import DynamicCache + +from .kv_caching import KeysValues + + +def kv2dc(cache: KeysValues) -> DynamicCache: + legacy_cache = tuple((kv_cache._k_cache.get(), kv_cache._v_cache.get()) for kv_cache in cache) + return DynamicCache.from_legacy_cache(legacy_cache) + + +def update_kv(cache: KeysValues, new_cache: DynamicCache) -> None: + for i, (key_cache, value_cache) in enumerate(new_cache.to_legacy_cache()): + cache[i].update(key_cache[:, :, -1:, :], value_cache[:, :, -1:, :]) + + +class HuggingfaceQwenTransformer(Qwen2ForCausalLM): + """Qwen2 backbone adapter exposing the minimal UniZero transformer interface.""" + + @classmethod + def from_pretrained(cls, lzero_config, *args, **kwargs): + model = super(HuggingfaceQwenTransformer, cls).from_pretrained(*args, **kwargs) + model.lzero_config = lzero_config + return model + + def generate_empty_keys_values(self, n: int, max_tokens: int) -> KeysValues: + device = torch.device(self.lzero_config.device) + if device.type == "cuda" and not torch.cuda.is_available(): + device = torch.device("cpu") + return KeysValues( + n, + self.lzero_config.num_heads, + max_tokens, + self.lzero_config.embed_dim, + self.lzero_config.num_layers, + device, + self.lzero_config.hidden_size, + ) + + def _get_positional_embedding(self, layer: int, attn_type: str, pos_emb) -> torch.Tensor: + if attn_type == 'key': + module_name = 'k_proj' + elif attn_type == 'value': + module_name = 'v_proj' + elif attn_type == 'query': + module_name = 'q_proj' + else: + raise ValueError(f"Unsupported attention projection type: {attn_type}") + attn_func = getattr(self.model.layers[layer].self_attn, module_name) + return attn_func(pos_emb.weight) + + def forward( + self, + sequences: torch.Tensor, + past_keys_values: Optional[KeysValues] = None, + valid_context_lengths: Optional[torch.Tensor] = None, + start_pos: int = 0, + ) -> torch.Tensor: + assert past_keys_values is None or len(past_keys_values) == len(self.model.layers) + if past_keys_values is not None: + kv_cache = kv2dc(past_keys_values) + use_cache = True + else: + kv_cache = None + use_cache = False + + batch_size, seq_len, _ = sequences.shape + if valid_context_lengths is not None: + position = torch.arange(seq_len, device=sequences.device).expand(batch_size, seq_len) + attention_mask = position >= (seq_len - valid_context_lengths.to(sequences.device).unsqueeze(1)) + else: + attention_mask = torch.ones(batch_size, seq_len, device=sequences.device, dtype=torch.long) + + output = self.model.forward( + attention_mask=attention_mask, + past_key_values=kv_cache, + inputs_embeds=sequences, + use_cache=use_cache, + ) + + if kv_cache is not None: + update_kv(past_keys_values, kv_cache) + return output.last_hidden_state diff --git a/lzero/model/unizero_world_models/kv_caching.py b/lzero/model/unizero_world_models/kv_caching.py index cf040b13a..af9f51c01 100644 --- a/lzero/model/unizero_world_models/kv_caching.py +++ b/lzero/model/unizero_world_models/kv_caching.py @@ -98,7 +98,15 @@ class Cache: in a Transformer-like model. It handles dynamic updates and size management. """ - def __init__(self, num_samples: int, num_heads: int, max_tokens: int, embed_dim: int, device: torch.device) -> None: + def __init__( + self, + num_samples: int, + num_heads: int, + max_tokens: int, + embed_dim: int, + device: torch.device, + hidden_size: Optional[int] = None, + ) -> None: """ Overview: Initializes the cache. @@ -115,7 +123,7 @@ def __init__(self, num_samples: int, num_heads: int, max_tokens: int, embed_dim: self._num_samples = num_samples self._num_heads = num_heads self._max_tokens = max_tokens - self._head_dim = embed_dim // num_heads + self._head_dim = hidden_size if hidden_size is not None else embed_dim // num_heads self._device = device self._cache: torch.Tensor = self._create_cache_tensor(self._num_samples) @@ -221,7 +229,15 @@ class KVCache: typically used in a single attention layer of a Transformer. """ - def __init__(self, num_samples: int, num_heads: int, max_tokens: int, embed_dim: int, device: torch.device) -> None: + def __init__( + self, + num_samples: int, + num_heads: int, + max_tokens: int, + embed_dim: int, + device: torch.device, + hidden_size: Optional[int] = None, + ) -> None: """ Overview: Initializes the Key-Value cache pair. @@ -232,8 +248,8 @@ def __init__(self, num_samples: int, num_heads: int, max_tokens: int, embed_dim: - embed_dim (:obj:`int`): The total dimension of the embeddings. - device (:obj:`torch.device`): The device on which to store the cache tensors. """ - self._k_cache = Cache(num_samples, num_heads, max_tokens, embed_dim, device) - self._v_cache = Cache(num_samples, num_heads, max_tokens, embed_dim, device) + self._k_cache = Cache(num_samples, num_heads, max_tokens, embed_dim, device, hidden_size) + self._v_cache = Cache(num_samples, num_heads, max_tokens, embed_dim, device, hidden_size) @property def shape(self) -> Tuple[int, int, int, int]: @@ -300,7 +316,8 @@ def __init__( max_tokens: int, embed_dim: int, num_layers: int, - device: torch.device + device: torch.device, + hidden_size: Optional[int] = None, ) -> None: """ Overview: @@ -314,7 +331,7 @@ def __init__( - device (:obj:`torch.device`): The device for storing cache tensors. """ self._keys_values = tuple([ - KVCache(num_samples, num_heads, max_tokens, embed_dim, device) for _ in range(num_layers) + KVCache(num_samples, num_heads, max_tokens, embed_dim, device, hidden_size) for _ in range(num_layers) ]) def __getitem__(self, layer_index: int) -> KVCache: @@ -384,4 +401,4 @@ def remove_register_tokens(self, register_token_num: int) -> None: for kv_cache in self._keys_values: # Decrement the size pointer for both K and V caches. kv_cache._k_cache._size = max(0, kv_cache._k_cache._size - register_token_num) - kv_cache._v_cache._size = max(0, kv_cache._v_cache._size - register_token_num) \ No newline at end of file + kv_cache._v_cache._size = max(0, kv_cache._v_cache._size - register_token_num) diff --git a/lzero/model/unizero_world_models/tokenizer.py b/lzero/model/unizero_world_models/tokenizer.py index 1035c46a7..309e86e82 100644 --- a/lzero/model/unizero_world_models/tokenizer.py +++ b/lzero/model/unizero_world_models/tokenizer.py @@ -144,6 +144,8 @@ def encode_to_obs_embeddings(self, x: torch.Tensor, task_id: int = 0) -> torch.T elif len(original_shape) == 3: # Batch of sequences of vectors: (B, T, E) # Flatten the batch and time dimensions to create a batch of vectors. x = x.contiguous().view(-1, original_shape[-1]) # Shape: (B*T, E) + elif len(original_shape) == 1: # Single observation without batch dim: (E,) + x = x.unsqueeze(0) # Shape: (1, E) # Note: 2D (B, E) and 4D (B, C, H, W) inputs are processed directly without reshaping. # [DEBUG] Log shape before encoder diff --git a/lzero/model/unizero_world_models/world_model.py b/lzero/model/unizero_world_models/world_model.py index d69671ac5..8ab0e1613 100644 --- a/lzero/model/unizero_world_models/world_model.py +++ b/lzero/model/unizero_world_models/world_model.py @@ -15,6 +15,7 @@ from .tokenizer import Tokenizer from .transformer import Transformer, TransformerConfig from .utils import LossWithIntermediateLosses, init_weights, WorldModelOutput, hash_state +from .hf_transformer import HuggingfaceQwenTransformer from collections import OrderedDict logging.getLogger().setLevel(logging.DEBUG) @@ -59,7 +60,13 @@ def __init__(self, config: TransformerConfig, tokenizer) -> None: self.config = config self.task_embed_option = self.config.task_embed_option # Strategy for task embeddings - self.transformer = Transformer(self.config) + if getattr(self.config, 'use_qwen_backbone', False): + self.transformer = HuggingfaceQwenTransformer.from_pretrained( + self.config, + self.config.pretrained_path, + ) + else: + self.transformer = Transformer(self.config) self.task_num = 1 self.env_num = self.config.env_num if self.config.device == 'cpu': @@ -78,7 +85,8 @@ def __init__(self, config: TransformerConfig, tokenizer) -> None: # Initialize patterns for block masks self._initialize_patterns() - self.hidden_size = config.embed_dim // config.num_heads + self.hidden_size = getattr(config, 'hidden_size', config.embed_dim // config.num_heads) + config['hidden_size'] = self.hidden_size # Position embedding if not self.config.rotary_emb: @@ -614,15 +622,20 @@ def _get_positional_embedding(self, layer, attn_type) -> torch.Tensor: Returns: - torch.Tensor: The positional embedding tensor. """ - attn_func = getattr(self.transformer.blocks[layer].attn, attn_type) - if torch.cuda.is_available(): - return attn_func(self.pos_emb.weight).view( - 1, self.config.max_tokens, self.num_heads, self.embed_dim // self.num_heads - ).transpose(1, 2).to(self.device).detach() + if getattr(self.config, 'use_qwen_backbone', False): + positional_embedding = self.transformer._get_positional_embedding(layer, attn_type, self.pos_emb) + positional_embedding = positional_embedding.view( + 1, self.config.max_tokens, self.num_heads, self.hidden_size + ) else: - return attn_func(self.pos_emb.weight).view( + attn_func = getattr(self.transformer.blocks[layer].attn, attn_type) + positional_embedding = attn_func(self.pos_emb.weight).view( 1, self.config.max_tokens, self.num_heads, self.embed_dim // self.num_heads - ).transpose(1, 2).detach() + ) + if torch.cuda.is_available(): + return positional_embedding.transpose(1, 2).to(self.device).detach() + else: + return positional_embedding.transpose(1, 2).detach() def forward( self, @@ -1636,6 +1649,8 @@ def retrieve_or_generate_kvcache(self, latent_state: list, ready_env_num: int, def compute_loss(self, batch, target_tokenizer: Tokenizer = None, inverse_scalar_transform_handle=None, **kwargs: Any) -> LossWithIntermediateLosses: + # import ipdb;ipdb.set_trace() + start_pos = batch['timestep'] # Encode observations into latent state representations @@ -2064,7 +2079,7 @@ def compute_loss(self, batch, target_tokenizer: Tokenizer = None, inverse_scalar value_priority=value_priority, intermediate_tensor_x=intermediate_tensor_x, obs_embeddings=detached_obs_embeddings, # <-- 新增 - ) + ), inverse_scalar_transform_handle(outputs.logits_value.reshape(-1, outputs.logits_value.shape[-1])).detach() # TODO: test correctness diff --git a/lzero/policy/unizero.py b/lzero/policy/unizero.py index 437817557..c764481ae 100644 --- a/lzero/policy/unizero.py +++ b/lzero/policy/unizero.py @@ -217,12 +217,12 @@ class UniZeroPolicy(MuZeroPolicy): ), # ****** common ****** # (bool) 是否启用自适应策略熵权重 (alpha) - use_adaptive_entropy_weight=True, + use_adaptive_entropy_weight=False, # (float) 自适应alpha优化器的学习率 adaptive_entropy_alpha_lr=1e-4, # ==================== START: Encoder-Clip Annealing Config ==================== # (bool) 是否启用 encoder-clip 值的退火。 - use_encoder_clip_annealing=True, + use_encoder_clip_annealing=False, # (str) 退火类型。可选 'linear' 或 'cosine'。 encoder_clip_anneal_type='cosine', # (float) 退火的起始 clip 值 (训练初期,较宽松)。 @@ -232,7 +232,7 @@ class UniZeroPolicy(MuZeroPolicy): # (int) 完成从起始值到结束值的退火所需的训练迭代步数。 encoder_clip_anneal_steps=100000, # 例如,在200k次迭代后达到最终值 # ===================== END: Encoder-Clip Annealing Config ===================== - + monitor_norm_freq=500000, # (bool) whether to use rnd model. use_rnd_model=False, # (bool) Whether to use multi-gpu training. @@ -654,8 +654,8 @@ def _forward_learn(self, data: Tuple[torch.Tensor]) -> Dict[str, Union[float, in # target_reward_categorical = phi_transform(self.reward_support, transformed_target_reward) # target_value_categorical = phi_transform(self.value_support, transformed_target_value) - target_reward_categorical = phi_transform(self.reward_support, transformed_target_reward, label_smoothing_eps= self._cfg.label_smoothing_eps) - target_value_categorical = phi_transform(self.value_support, transformed_target_value, label_smoothing_eps=self._cfg.label_smoothing_eps) + target_reward_categorical = phi_transform(self.reward_support, transformed_target_reward) + target_value_categorical = phi_transform(self.value_support, transformed_target_value) # Prepare batch for GPT model batch_for_gpt = {} @@ -686,42 +686,10 @@ def _forward_learn(self, data: Tuple[torch.Tensor]) -> Dict[str, Union[float, in average_target_policy_entropy = target_policy_entropy.mean() # Update world model - losses = self._learn_model.world_model.compute_loss( + losses, _ = self._learn_model.world_model.compute_loss( batch_for_gpt, self._target_model.world_model.tokenizer, self.value_inverse_scalar_transform_handle, global_step=train_iter, current_policy_label_eps=current_policy_label_eps, ) # NOTE : compute_loss third argument is now a dead argument. If this changes, it could need adaptation between value_inverse and reward_inverse. - # ==================== [修改] 集成范数监控逻辑 ==================== - norm_log_dict = {} - # 检查是否达到监控频率 - if self._cfg.monitor_norm_freq > 0 and train_iter == 0 or (train_iter % self._cfg.monitor_norm_freq == 0): - with torch.no_grad(): - # 1. 监控模型参数范数 - param_norm_metrics = self._monitor_model_norms() - norm_log_dict.update(param_norm_metrics) - - # 2. 监控中间张量 x (Transformer的输出) - intermediate_x = losses.intermediate_losses.get('intermediate_tensor_x') - if intermediate_x is not None: - # x 的形状为 (B, T, E) - # 计算每个 token 的 L2 范数 - token_norms = intermediate_x.norm(p=2, dim=-1) - - # 记录这些范数的统计数据 - norm_log_dict['norm/x_token/mean'] = token_norms.mean().item() - norm_log_dict['norm/x_token/std'] = token_norms.std().item() - norm_log_dict['norm/x_token/max'] = token_norms.max().item() - norm_log_dict['norm/x_token/min'] = token_norms.min().item() - # ================================================================= - - # ==================== START MODIFICATION 2 ==================== - # Extract the calculated value_priority from the returned losses. - value_priority_tensor = losses.intermediate_losses['value_priority'] - # Convert to numpy array for the replay buffer, adding a small epsilon. - value_priority_np = value_priority_tensor.detach().cpu().numpy() + 1e-6 - # ===================== END MODIFICATION 2 ===================== - - # weighted_total_loss = losses.loss_total - # TODO: weighted_total_loss = (weights * losses.loss_total).mean() for loss_name, loss_value in losses.intermediate_losses.items(): @@ -761,63 +729,53 @@ def _forward_learn(self, data: Tuple[torch.Tensor]) -> Dict[str, Union[float, in temperature_reward=self.intermediate_losses['temperature_reward'] temperature_policy=self.intermediate_losses['temperature_policy'] - assert not torch.isnan(losses.loss_total).any(), "Loss contains NaN values" - assert not torch.isinf(losses.loss_total).any(), "Loss contains Inf values" - - # Core learning model update step - # Reset gradients at the start of each accumulation cycle - if (train_iter % self.accumulation_steps) == 0: - self._optimizer_world_model.zero_grad() - - # ==================== START: 目标熵正则化更新逻辑 ==================== - alpha_loss = None - current_alpha = self._cfg.model.world_model_cfg.policy_entropy_weight # 默认使用固定值 + current_alpha = self._cfg.model.world_model_cfg.policy_entropy_weight # 默认使用固定值 + current_ratio = 0.0 + alpha_loss = torch.tensor(0.0, device=self._cfg.device) if self.use_adaptive_entropy_weight: - # --- 动态计算目标熵 (这部分逻辑是正确的,予以保留) --- + # --- 动态计算目标熵 --- progress = min(1.0, train_iter / self.target_entropy_decay_steps) current_ratio = self.target_entropy_start_ratio * (1 - progress) + self.target_entropy_end_ratio * progress action_space_size = self._cfg.model.action_space_size - # 注意:我们将 target_entropy 定义为正数,更符合直觉 current_target_entropy = -np.log(1.0 / action_space_size) * current_ratio - # --- 计算 alpha_loss (已修正符号) --- - # 这是核心修正点:去掉了最前面的负号 - # detach() 仍然是关键,确保 alpha_loss 的梯度只流向 log_alpha + # --- 计算 alpha_loss --- alpha_loss = (self.log_alpha * (policy_entropy.detach() - current_target_entropy)).mean() - # # --- 更新 log_alpha --- + # --- 更新 log_alpha --- self.alpha_optimizer.zero_grad() alpha_loss.backward() self.alpha_optimizer.step() - # --- [优化建议] 增加 log_alpha 裁剪作为安全措施 --- with torch.no_grad(): - # 将 alpha 限制在例如 [1e-4, 10.0] 的范围内 self.log_alpha.clamp_(np.log(1e-4), np.log(10.0)) # --- 使用当前更新后的 alpha (截断梯度流) --- current_alpha = self.log_alpha.exp().detach() # 重新计算加权的策略损失和总损失 - # 注意:这里的 policy_entropy 已经是一个batch的平均值 weighted_policy_loss = orig_policy_loss - current_alpha * policy_entropy - # 重新构建总损失 (不使用 losses.loss_total) - # 确保这里的权重与 LossWithIntermediateLosses 类中的计算方式一致 self.obs_loss_weight = 10 self.value_loss_weight = 0.5 self.reward_loss_weight = 1. self.policy_loss_weight = 1. - self.ends_loss_weight = 0. total_loss = ( self.reward_loss_weight * reward_loss + self.value_loss_weight * value_loss + self.policy_loss_weight * weighted_policy_loss + - self.obs_loss_weight * obs_loss # 假设 ssl_loss_weight 是 obs_loss 的权重 - # ... 如果还有其他损失项,也加进来 ... + self.obs_loss_weight * obs_loss ) weighted_total_loss = (weights * total_loss).mean() # ===================== END: 目标熵正则化更新逻辑 ===================== + assert not torch.isnan(losses.loss_total).any(), "Loss contains NaN values" + assert not torch.isinf(losses.loss_total).any(), "Loss contains Inf values" + + # Core learning model update step + # Reset gradients at the start of each accumulation cycle + if (train_iter % self.accumulation_steps) == 0: + self._optimizer_world_model.zero_grad() + # Scale the loss by the number of accumulation steps weighted_total_loss = weighted_total_loss / self.accumulation_steps weighted_total_loss.backward() @@ -843,7 +801,7 @@ def _forward_learn(self, data: Tuple[torch.Tensor]) -> Dict[str, Union[float, in # ===================== END: 动态计算当前 Clip 阈值 ===================== # 1. Encoder-Clip (使用动态计算出的 current_clip_value) - if current_clip_value > 0 and 'obs_embeddings' in losses.intermediate_losses: + if self.use_encoder_clip_annealing and current_clip_value > 0 and 'obs_embeddings' in losses.intermediate_losses: obs_embeddings = losses.intermediate_losses['obs_embeddings'] if obs_embeddings is not None: max_latent_norm = obs_embeddings.norm(p=2, dim=-1).max() @@ -930,8 +888,6 @@ def _forward_learn(self, data: Tuple[torch.Tensor]) -> Dict[str, Union[float, in 'reward_loss': reward_loss.item(), 'value_loss': value_loss.item(), # Add value_priority to the log dictionary. - 'value_priority': value_priority_np.mean().item(), - 'value_priority_orig': value_priority_np, 'target_reward': target_reward.mean().item(), 'target_value': target_value.mean().item(), 'transformed_target_reward': transformed_target_reward.mean().item(), @@ -979,6 +935,12 @@ def _forward_learn(self, data: Tuple[torch.Tensor]) -> Dict[str, Union[float, in return_log_dict['current_encoder_clip_value'] = current_clip_value # ===================== END: 添加新日志项 ===================== + if getattr(self._cfg.model.world_model_cfg, 'use_qwen_backbone', False): + for key, value in list(return_log_dict.items()): + if isinstance(value, torch.Tensor): + value = value.detach() + return_log_dict[key] = value.item() if value.numel() == 1 else value.float().mean().item() + if self._cfg.use_wandb: wandb.log({'learner_step/' + k: v for k, v in return_log_dict.items()}, step=self.env_step) wandb.log({"learner_iter_vs_env_step": self.train_iter}, step=self.env_step) @@ -1011,13 +973,13 @@ def _init_collect(self) -> None: self._collect_epsilon = 0.0 self.collector_env_num = self._cfg.collector_env_num if self._cfg.model.model_type == 'conv': - self.last_batch_obs = torch.zeros([self.collector_env_num, self._cfg.model.observation_shape[0], 64, 64]).to(self._cfg.device) - self.last_batch_action = [-1 for i in range(self.collector_env_num)] + self.last_batch_obs_collect = torch.zeros([self.collector_env_num, self._cfg.model.observation_shape[0], 64, 64]).to(self._cfg.device) + self.last_batch_action_collect = [-1 for i in range(self.collector_env_num)] elif self._cfg.model.model_type == 'mlp': - self.last_batch_obs = torch.full( + self.last_batch_obs_collect = torch.full( [self.collector_env_num, self._cfg.model.observation_shape], fill_value=self.pad_token_id, ).to(self._cfg.device) - self.last_batch_action = [-1 for i in range(self.collector_env_num)] + self.last_batch_action_collect = [-1 for i in range(self.collector_env_num)] # @profile def _forward_collect( @@ -1067,7 +1029,7 @@ def _forward_collect( output = {i: None for i in ready_env_id} with torch.no_grad(): - network_output = self._collect_model.initial_inference(self.last_batch_obs, self.last_batch_action, data, timestep) + network_output = self._collect_model.initial_inference(self.last_batch_obs_collect, self.last_batch_action_collect, data, timestep) latent_state_roots, reward_roots, pred_values, policy_logits = mz_network_output_unpack(network_output) pred_values = self.value_inverse_scalar_transform_handle(pred_values).detach().cpu().numpy() @@ -1145,18 +1107,18 @@ def _forward_collect( } batch_action.append(action) - self.last_batch_obs = data - self.last_batch_action = batch_action + self.last_batch_obs_collect = data + self.last_batch_action_collect = batch_action # ========= TODO: This logic is a temporary workaround specific to the muzero_segment_collector. ========= if active_collect_env_num < self.collector_env_num: - # When an environment finishes an episode ('done'), the length of `self.last_batch_obs` passed back + # When an environment finishes an episode ('done'), the length of `self.last_batch_obs_collect` passed back # becomes smaller than the total number of collector environments. # Handling this dynamic batch size is complex, as the transformer's KV cache retrieval # requires a stable environment ID for correct indexing. A mismatch would cause retrieval errors. # # Therefore, as a simpler solution, we reset the collection state for ALL environments. - # By resetting `self.last_batch_action` to -1 for all `self.collector_env_num` environments, + # By resetting `self.last_batch_action_collect` to -1 for all `self.collector_env_num` environments, # we force the transformer to start its context from scratch, avoiding incorrect cache lookups. print('========== collect_forward ============') print(f'An environment has finished. Active envs: {active_collect_env_num} < Total envs: {self.collector_env_num}. Resetting all.') @@ -1189,13 +1151,13 @@ def _init_eval(self) -> None: self.evaluator_env_num = self._cfg.evaluator_env_num if self._cfg.model.model_type == 'conv': - self.last_batch_obs = torch.zeros([self.collector_env_num, self._cfg.model.observation_shape[0], 64, 64]).to(self._cfg.device) - self.last_batch_action = [-1 for i in range(self.collector_env_num)] + self.last_batch_obs_eval = torch.zeros([self.evaluator_env_num, self._cfg.model.observation_shape[0], 64, 64]).to(self._cfg.device) + self.last_batch_action_eval = [-1 for i in range(self.evaluator_env_num)] elif self._cfg.model.model_type == 'mlp': - self.last_batch_obs = torch.full( - [self.collector_env_num, self._cfg.model.observation_shape], fill_value=self.pad_token_id, + self.last_batch_obs_eval = torch.full( + [self.evaluator_env_num, self._cfg.model.observation_shape], fill_value=self.pad_token_id, ).to(self._cfg.device) - self.last_batch_action = [-1 for i in range(self.collector_env_num)] + self.last_batch_action_eval = [-1 for i in range(self.evaluator_env_num)] def _forward_eval(self, data: torch.Tensor, action_mask: list, to_play: int = -1, ready_env_id: np.array = None, timestep: List = [0], task_id: int = None,) -> Dict: @@ -1226,11 +1188,13 @@ def _forward_eval(self, data: torch.Tensor, action_mask: list, to_play: int = -1 """ self._eval_model.eval() active_eval_env_num = data.shape[0] + if active_eval_env_num == 0 or data.numel() == 0: + return {} if ready_env_id is None: ready_env_id = np.arange(active_eval_env_num) output = {i: None for i in ready_env_id} with torch.no_grad(): - network_output = self._eval_model.initial_inference(self.last_batch_obs_eval, self.last_batch_action, data, timestep) + network_output = self._eval_model.initial_inference(self.last_batch_obs_eval, self.last_batch_action_eval, data, timestep) latent_state_roots, reward_roots, pred_values, policy_logits = mz_network_output_unpack(network_output) # if not in training, obtain the scalars of the value/reward @@ -1291,7 +1255,7 @@ def _forward_eval(self, data: torch.Tensor, action_mask: list, to_play: int = -1 batch_action.append(action) self.last_batch_obs_eval = data - self.last_batch_action = batch_action + self.last_batch_action_eval = batch_action return output @@ -1308,13 +1272,13 @@ def _reset_collect(self, env_id: int = None, current_steps: int = None, reset_in - reset_init_data (:obj:`bool`, optional): Whether to reset the initial data. If True, the initial data will be reset. """ if reset_init_data: - self.last_batch_obs = initialize_pad_batch( + self.last_batch_obs_collect = initialize_pad_batch( self._cfg.model.observation_shape, self._cfg.collector_env_num, self._cfg.device, pad_token_id=self.pad_token_id ) - self.last_batch_action = [-1 for _ in range(self._cfg.collector_env_num)] + self.last_batch_action_collect = [-1 for _ in range(self._cfg.collector_env_num)] # We must handle both single int and list of ints for env_id. @@ -1386,7 +1350,7 @@ def _reset_eval(self, env_id: int = None, current_steps: int = None, reset_init_ ) print(f'unizero.py task_id:{task_id} after _reset_eval: last_batch_obs_eval:', self.last_batch_obs_eval.shape) - self.last_batch_action = [-1 for _ in range(self._cfg.evaluator_env_num)] + self.last_batch_action_eval = [-1 for _ in range(self._cfg.evaluator_env_num)] # --- BEGIN ROBUST FIX --- # This logic handles the crucial end-of-episode cache clearing for evaluation. @@ -1587,4 +1551,4 @@ def recompute_pos_emb_diff_and_clear_cache(self) -> None: # If rotary_emb is False, nn.Embedding is used for absolute position encoding. model.world_model.precompute_pos_emb_diff_kv() model.world_model.clear_caches() - torch.cuda.empty_cache() \ No newline at end of file + torch.cuda.empty_cache() diff --git a/lzero/policy/unizero_multitask_alpha_indep.py b/lzero/policy/unizero_multitask_alpha_indep.py deleted file mode 100644 index db2b4c513..000000000 --- a/lzero/policy/unizero_multitask_alpha_indep.py +++ /dev/null @@ -1,2000 +0,0 @@ -import copy -from collections import defaultdict -from typing import List, Dict, Any, Tuple, Union - -import numpy as np -import torch -from ding.model import model_wrap -from ding.utils import POLICY_REGISTRY - -from lzero.entry.utils import initialize_zeros_batch -from lzero.mcts import UniZeroMCTSCtree as MCTSCtree -from lzero.model import ImageTransforms -from lzero.policy import prepare_obs_stack_for_unizero -from lzero.policy import scalar_transform, InverseScalarTransform, phi_transform, \ - DiscreteSupport, to_torch_float_tensor, mz_network_output_unpack, select_action, prepare_obs -from lzero.policy.unizero import UniZeroPolicy, scale_module_weights_vectorized -from .utils import configure_optimizers_nanogpt -import sys - -# Please replace the path with the actual location of your LibMTL library. -sys.path.append('/path/to/your/LibMTL') - -from LibMTL.weighting.MoCo_unizero import MoCo as GradCorrect -from LibMTL.weighting.moco_fast_mem_eff import FastMoCoMemEff as FastMoCo -from LibMTL.weighting.moco_fast_mem_eff import MoCoCfg - -import torch.distributed as dist - -# ------------------------------------------------------------ -# 1. Add a dedicated process-group for the learner. -# (This function should be called once during the initialization of the main process or the learner.) -# ------------------------------------------------------------ -def build_learner_group(learner_ranks: list[int]) -> dist.ProcessGroup: - """ - Overview: - Builds and returns a new process group containing only the learner ranks. - This is used for methods like GenericMoCo that require collective communication - only among the ranks performing training. - Arguments: - - learner_ranks (:obj:`list[int]`): A list of world ranks that are designated as learners. - These are the ranks that will perform the backward pass. - e.g., if CUDA_VISIBLE_DEVICES=0,1, then learner_ranks=[0,1]. - Returns: - - pg (:obj:`dist.ProcessGroup`): A new process group containing only the learner ranks. - """ - world_pg = dist.group.WORLD - pg = dist.new_group(ranks=learner_ranks, backend='nccl') - if dist.get_rank() in learner_ranks: - torch.cuda.set_device(learner_ranks.index(dist.get_rank())) - return pg - - -def generate_task_loss_dict(multi_task_losses: List[Union[torch.Tensor, float]], task_name_template: str, task_id: int) -> Dict[str, float]: - """ - Overview: - Generates a dictionary for the losses of each task. - Arguments: - - multi_task_losses (:obj:`List[Union[torch.Tensor, float]]`): A list containing the loss for each task. - - task_name_template (:obj:`str`): The template for the task name, e.g., 'obs_loss_task{}'. - - task_id (:obj:`int`): The starting ID of the tasks. - Returns: - - task_loss_dict (:obj:`Dict[str, float]`): A dictionary where keys are formatted task names and values are the corresponding losses. - """ - task_loss_dict = {} - for task_idx, task_loss in enumerate(multi_task_losses): - task_name = task_name_template.format(task_idx + task_id) - try: - # Get the scalar value of the loss if it's a tensor. - task_loss_dict[task_name] = task_loss.item() if hasattr(task_loss, 'item') else task_loss - except Exception as e: - task_loss_dict[task_name] = task_loss - return task_loss_dict - -# # 修改后的函数: -# def generate_task_loss_dict( -# multi_task_losses: List[Union[torch.Tensor, float]], -# task_name_template: str, -# global_task_ids: List[int] -# ) -> Dict[str, float]: -# """ -# Overview: -# Generates a dictionary for the losses of each task using their explicit global IDs. -# Arguments: -# - multi_task_losses (:obj:`List[Union[torch.Tensor, float]]`): A list containing the loss for each task. -# - task_name_template (:obj:`str`): The template for the task name, e.g., 'obs_loss_task{}'. -# - global_task_ids (:obj:`List[int]`): A list of global task IDs corresponding to each loss in multi_task_losses. -# Returns: -# - task_loss_dict (:obj:`Dict[str, float]`): A dictionary where keys are formatted task names and values are the corresponding losses. -# """ -# task_loss_dict = {} -# # 使用 zip 将每个损失与其正确的全局ID配对 -# for task_loss, global_id in zip(multi_task_losses, global_task_ids): -# task_name = task_name_template.format(global_id) -# try: -# task_loss_dict[task_name] = task_loss.item() if hasattr(task_loss, 'item') else task_loss -# except Exception as e: -# task_loss_dict[task_name] = task_loss -# return task_loss_dict - - -class WrappedModel: - """ - Overview: - A wrapper class for the world model to conveniently access its parameters and zero its gradients. - This version wraps the entire world model. - """ - def __init__(self, world_model: torch.nn.Module): - """ - Arguments: - - world_model (:obj:`torch.nn.Module`): The world model instance. - """ - self.world_model = world_model - - def parameters(self) -> iter: - """ - Overview: - Returns an iterator over the parameters of the entire world model. - """ - return self.world_model.parameters() - - def zero_grad(self, set_to_none: bool = False) -> None: - """ - Overview: - Sets the gradients of all world model parameters to zero. - Arguments: - - set_to_none (:obj:`bool`): Whether to set gradients to None instead of zero. - """ - self.world_model.zero_grad(set_to_none=set_to_none) - - -class WrappedModelV2: - """ - Overview: - A wrapper for specific components of the world model. - This version is designed to group parameters that are considered "shared" - across tasks for gradient correction methods like MoCo, excluding the prediction heads. - """ - def __init__(self, tokenizer: torch.nn.Module, transformer: torch.nn.Module, pos_emb: torch.nn.Module, task_emb: torch.nn.Module, act_embedding_table: torch.nn.Module): - """ - Arguments: - - tokenizer (:obj:`torch.nn.Module`): The tokenizer module. - - transformer (:obj:`torch.nn.Module`): The transformer backbone. - - pos_emb (:obj:`torch.nn.Module`): The positional embedding module. - - task_emb (:obj:`torch.nn.Module`): The task embedding module. - - act_embedding_table (:obj:`torch.nn.Module`): The action embedding table. - """ - self.tokenizer = tokenizer - self.transformer = transformer - self.pos_emb = pos_emb - self.task_emb = task_emb - self.act_embedding_table = act_embedding_table - - def parameters(self) -> iter: - """ - Overview: - Returns an iterator over the parameters of the wrapped components (tokenizer, transformer, embeddings). - These are typically the shared parts of the model whose gradients need to be managed for multi-task learning. - """ - return (list(self.tokenizer.parameters()) + - list(self.transformer.parameters()) + - list(self.pos_emb.parameters()) + - # list(self.task_emb.parameters()) + # TODO: Decide whether to include task embeddings in shared parameters. - list(self.act_embedding_table.parameters())) - - def zero_grad(self, set_to_none: bool = False) -> None: - """ - Overview: - Sets the gradients of all wrapped components to zero. - Arguments: - - set_to_none (:obj:`bool`): Whether to set gradients to None instead of zero. - """ - self.tokenizer.zero_grad(set_to_none=set_to_none) - self.transformer.zero_grad(set_to_none=set_to_none) - self.pos_emb.zero_grad(set_to_none=set_to_none) - # self.task_emb.zero_grad(set_to_none=set_to_none) # TODO: Match the decision made in the parameters() method. - self.act_embedding_table.zero_grad(set_to_none=set_to_none) - - -class WrappedModelV3: - """ - Overview: - An alternative wrapper for world model components. - This version excludes the tokenizer from the shared parameters, focusing gradient correction - on the transformer and embedding layers. - """ - def __init__(self, transformer: torch.nn.Module, pos_emb: torch.nn.Module, task_emb: torch.nn.Module, act_embedding_table: torch.nn.Module): - """ - Arguments: - - transformer (:obj:`torch.nn.Module`): The transformer backbone. - - pos_emb (:obj:`torch.nn.Module`): The positional embedding module. - - task_emb (:obj:`torch.nn.Module`): The task embedding module. - - act_embedding_table (:obj:`torch.nn.Module`): The action embedding table. - """ - self.transformer = transformer - self.pos_emb = pos_emb - self.task_emb = task_emb - self.act_embedding_table = act_embedding_table - - def parameters(self) -> iter: - """ - Overview: - Returns an iterator over the parameters of the transformer and various embedding layers. - """ - return (list(self.transformer.parameters()) + - list(self.pos_emb.parameters()) + - list(self.task_emb.parameters()) + - list(self.act_embedding_table.parameters())) - - def zero_grad(self, set_to_none: bool = False) -> None: - """ - Overview: - Sets the gradients of the wrapped components to zero. - Arguments: - - set_to_none (:obj:`bool`): Whether to set gradients to None instead of zero. - """ - self.transformer.zero_grad(set_to_none=set_to_none) - self.pos_emb.zero_grad(set_to_none=set_to_none) - self.task_emb.zero_grad(set_to_none=set_to_none) - self.act_embedding_table.zero_grad(set_to_none=set_to_none) - - -# def configure_optimizer_unizero(model, learning_rate, weight_decay, device_type, betas): -# """ -# 为UniZero模型配置带有差异化学习率的优化器。 -# """ -# # 1. 定义需要特殊处理的参数 -# param_dict = {pn: p for pn, p in model.named_parameters() if p.requires_grad} - -# # 2. 将参数分为三组:Transformer主干、Tokenizer、Heads -# transformer_params = {pn: p for pn, p in param_dict.items() if 'transformer' in pn} -# tokenizer_params = {pn: p for pn, p in param_dict.items() if 'tokenizer' in pn} - -# # Heads的参数是那些既不属于transformer也不属于tokenizer的 -# head_params = { -# pn: p for pn, p in param_dict.items() -# if 'transformer' not in pn and 'tokenizer' not in pn -# } - -# # 3. 为每组设置不同的优化器参数(特别是学习率) -# # 这里我们仍然使用AdamW,但学习率设置更合理 -# optim_groups = [ -# { -# 'params': list(transformer_params.values()), -# 'lr': learning_rate, # 1e-4 -# # 'lr': learning_rate * 0.2, # 为Transformer主干设置一个较小的学习率,例如 1e-5 -# 'weight_decay': weight_decay -# # 'weight_decay': weight_decay * 5.0 -# }, -# { -# 'params': list(tokenizer_params.values()), -# 'lr': learning_rate, # Tokenizer使用基础学习率,例如 1e-4 -# # 'lr': learning_rate * 0.1, # 为encoder设置一个较小的学习率,例如 1e-5 -# 'weight_decay': weight_decay * 5.0 # <-- 为Encoder设置5倍的权重衰减!这是一个强力正则化 - -# }, -# { -# 'params': list(head_params.values()), -# 'lr': learning_rate, # Heads也使用基础学习率率,例如 1e-4 -# 'weight_decay': 0.0 # 通常Heads的权重不做衰减 -# # 'weight_decay': weight_decay - -# } -# ] - -# print("--- Optimizer Groups ---") -# print(f"Transformer LR: {learning_rate}") -# print(f"Tokenizer/Heads LR: {learning_rate}") - -# optimizer = torch.optim.AdamW(optim_groups, betas=betas) -# return optimizer - -def configure_optimizer_unizero(model, learning_rate, weight_decay, device_type, betas): - """ - 为UniZero模型配置带有差异化学习率的优化器。 - (修正版,确保参数组互斥) - """ - # 1. 创建空的参数列表用于分组 - transformer_params = [] - tokenizer_params = [] - head_params = [] - - # 2. 遍历所有可训练参数,并使用 if/elif/else 结构确保每个参数只被分配到一个组 - for name, param in model.named_parameters(): - if not param.requires_grad: - continue - - if 'transformer' in name: - transformer_params.append(param) - elif 'tokenizer' in name: - tokenizer_params.append(param) - else: - head_params.append(param) - - # 3. 为每组设置不同的优化器参数 - # 这里我们仍然使用AdamW,但学习率设置更合理 - optim_groups = [ - { - 'params': transformer_params, - 'lr': learning_rate, # 1e-4 - 'weight_decay': weight_decay - }, - { - 'params': tokenizer_params, - 'lr': learning_rate, # Tokenizer使用基础学习率,例如 1e-4 - # 'weight_decay': weight_decay * 5.0 # <-- 为Encoder设置5倍的权重衰减!这是一个强力正则化 - 'weight_decay': weight_decay # <-- 为Encoder设置5倍的权重衰减!这是一个强力正则化 - }, - { - 'params': head_params, - 'lr': learning_rate, # Heads也使用基础学习率率,例如 1e-4 - # 'weight_decay': 0.0 # 通常Heads的权重不做衰减 - 'weight_decay': weight_decay - - } - ] - - print("--- Optimizer Groups ---") - # 打印每个组的参数数量以供调试 - print(f"Transformer params: {len(transformer_params)}") - print(f"Tokenizer params: {len(tokenizer_params)}") - print(f"Head params: {len(head_params)}") - print(f"Transformer LR: {learning_rate}") - print(f"Tokenizer/Heads LR: {learning_rate}") - - optimizer = torch.optim.AdamW(optim_groups, betas=betas) - return optimizer - -@POLICY_REGISTRY.register('unizero_multitask') -class UniZeroMTPolicy(UniZeroPolicy): - """ - Overview: - The policy class for multi-task UniZero, an official implementation for the paper "UniZero: Generalized and Efficient Planning - with Scalable Latent World Models". UniZero aims to enhance the planning capabilities of reinforcement learning agents - by addressing the limitations of MuZero-style algorithms, particularly in environments requiring the - capture of long-term dependencies. More details can be found at: https://arxiv.org/abs/2406.10667. - """ - - # The default_config for UniZero multi-task policy. - config = dict( - type='unizero_multitask', - model=dict( - # (str) The model type. For 1-dimensional vector obs, we use mlp model. For the image obs, we use conv model. - model_type='conv', # options={'mlp', 'conv'} - # (bool) If True, the action space of the environment is continuous, otherwise discrete. - continuous_action_space=False, - # (tuple) The obs shape. - observation_shape=(3, 64, 64), - # (bool) Whether to use the self-supervised learning loss. - self_supervised_learning_loss=True, - # (bool) Whether to use discrete support to represent categorical distribution for value/reward/value_prefix. - categorical_distribution=True, - # (int) The image channel in image observation. - image_channel=3, - # (int) The number of frames to stack together. - frame_stack_num=1, - # (int) The number of res blocks in MuZero model. - num_res_blocks=1, - # (int) The number of channels of hidden states in MuZero model. - num_channels=64, - # (int) The scale of supports used in categorical distribution. - # This variable is only effective when ``categorical_distribution=True``. - support_scale=50, - # (bool) whether to learn bias in the last linear layer in value and policy head. - bias=True, - # (bool) whether to use res connection in dynamics. - res_connection_in_dynamics=True, - # (str) The type of normalization in MuZero model. Options are ['BN', 'LN']. Default to 'BN'. - norm_type='LN', # NOTE: LayerNorm is used in the transformer-based world model. - # (bool) Whether to analyze simulation normalization. - analysis_sim_norm=False, - # (int) The save interval of the model. - learn=dict(learner=dict(hook=dict(save_ckpt_after_iter=10000, ), ), ), - world_model_cfg=dict( - # (int) The number of tokens per block. - tokens_per_block=2, - # (int) The maximum number of blocks. - max_blocks=10, - # (int) The maximum number of tokens, calculated as tokens per block multiplied by max blocks. - max_tokens=2 * 10, - # (int) The context length, usually calculated as twice the number of some base unit. - context_length=2 * 4, - # (bool) Whether to use GRU gating mechanism. - gru_gating=False, - # (str) The device to be used for computation, e.g., 'cpu' or 'cuda'. - device='cpu', - # (bool) Whether to analyze simulation normalization. - analysis_sim_norm=False, - # (bool) Whether to analyze dormant ratio. - analysis_dormant_ratio=False, - # (int) The shape of the action space. - action_space_size=6, - # (int) The size of the group, related to simulation normalization. - group_size=8, # NOTE: for sim_norm - # (str) The type of attention mechanism used. Options could be ['causal']. - attention='causal', - # (int) The number of layers in the model. - num_layers=2, - # (int) The number of attention heads. - num_heads=8, - # (int) The dimension of the embedding. - embed_dim=768, - # (float) The dropout probability for the embedding layer. - embed_pdrop=0.1, - # (float) The dropout probability for the residual connections. - resid_pdrop=0.1, - # (float) The dropout probability for the attention mechanism. - attn_pdrop=0.1, - # (int) The size of the support set for value and reward heads. - support_size=101, - # (int) The maximum size of the cache. - max_cache_size=5000, - # (int) The number of environments. - env_num=8, - # (float) The weight of the latent reconstruction loss. - latent_recon_loss_weight=0., - # (float) The weight of the perceptual loss. - perceptual_loss_weight=0., - # (float) The weight of the policy entropy. - policy_entropy_weight=1e-4, - # (str) The type of loss for predicting latent variables. Options could be ['group_kl', 'mse']. - predict_latent_loss_type='group_kl', - # (str) The type of observation. Options are ['image', 'vector']. - obs_type='image', - # (float) The discount factor for future rewards. - gamma=1, - # (bool) Whether to analyze dormant ratio, average_weight_magnitude of net, effective_rank of latent. - analysis_dormant_ratio_weight_rank=False, - # (float) The threshold for a dormant neuron. - dormant_threshold=0.01, - - ), - ), - # ****** common ****** - # (bool) whether to use rnd model. - use_rnd_model=False, - # (bool) Whether to use multi-gpu training. - multi_gpu=True, - # (bool) Whether to enable the sampled-based algorithm (e.g. Sampled EfficientZero) - # this variable is used in ``collector``. - sampled_algo=False, - # (bool) Whether to enable the gumbel-based algorithm (e.g. Gumbel Muzero) - gumbel_algo=False, - # (bool) Whether to use C++ MCTS in policy. If False, use Python implementation. - mcts_ctree=True, - # (bool) Whether to use cuda for network. - cuda=True, - # (int) The number of environments used in collecting data. - collector_env_num=8, - # (int) The number of environments used in evaluating policy. - evaluator_env_num=3, - # (str) The type of environment. Options are ['not_board_games', 'board_games']. - env_type='not_board_games', - # (str) The type of action space. Options are ['fixed_action_space', 'varied_action_space']. - action_type='fixed_action_space', - # (str) The type of battle mode. Options are ['play_with_bot_mode', 'self_play_mode']. - battle_mode='play_with_bot_mode', - # (bool) Whether to monitor extra statistics in tensorboard. - monitor_extra_statistics=True, - # (int) The transition number of one ``GameSegment``. - game_segment_length=400, - # (bool) Whether to analyze simulation normalization. - analysis_sim_norm=False, - # (bool) Whether to use the pure policy to collect data. - collect_with_pure_policy=False, - # (int) The evaluation frequency. - eval_freq=int(5e3), - # (str) The sample type. Options are ['episode', 'transition']. - sample_type='transition', - - # ****** observation ****** - # (bool) Whether to transform image to string to save memory. - transform2string=False, - # (bool) Whether to use gray scale image. - gray_scale=False, - # (bool) Whether to use data augmentation. - use_augmentation=False, - # (list) The style of augmentation. - augmentation=['shift', 'intensity'], - - # ******* learn ****** - # (bool) Whether to ignore the done flag in the training data. Typically, this value is set to False. - # However, for some environments with a fixed episode length, to ensure the accuracy of Q-value calculations, - # we should set it to True to avoid the influence of the done flag. - ignore_done=False, - # (int) How many updates(iterations) to train after collector's one collection. - # Bigger "update_per_collect" means bigger off-policy. - # collect data -> update policy-> collect data -> ... - # For different env, we have different episode_length, - # we usually set update_per_collect = collector_env_num * episode_length / batch_size * reuse_factor. - # If we set update_per_collect=None, we will set update_per_collect = collected_transitions_num * cfg.policy.replay_ratio automatically. - update_per_collect=None, - # (float) The ratio of the collected data used for training. Only effective when ``update_per_collect`` is not None. - replay_ratio=0.25, - # (int) Minibatch size for one gradient descent. - batch_size=256, - # (str) Optimizer for training policy network. - optim_type='AdamW', - # (float) Learning rate for training policy network. Initial lr for manually decay schedule. - learning_rate=0.0001, - # (int) Frequency of hard target network update. - target_update_freq=100, - # (int) Frequency of soft target network update. - target_update_theta=0.05, - # (int) Frequency of target network update. - target_update_freq_for_intrinsic_reward=1000, - # (float) Weight decay for training policy network. - weight_decay=1e-4, - # (float) One-order Momentum in optimizer, which stabilizes the training process (gradient direction). - momentum=0.9, - # (float) The maximum constraint value of gradient norm clipping. - grad_clip_value=5, - # (int) The number of episodes in each collecting stage when use muzero_collector. - n_episode=8, - # (int) The number of num_segments in each collecting stage when use muzero_segment_collector. - num_segments=8, - # # (int) the number of simulations in MCTS for renalyze. - num_simulations=50, - # (int) The number of simulations in MCTS for the collect phase. - collect_num_simulations=25, - # (int) The number of simulations in MCTS for the eval phase. - eval_num_simulations=50, - # (float) Discount factor (gamma) for returns. - discount_factor=0.997, - # (int) The number of steps for calculating target q_value. - td_steps=5, - # (int) The number of unroll steps in dynamics network. - num_unroll_steps=10, - # (float) The weight of reward loss. - reward_loss_weight=1, - # (float) The weight of value loss. - value_loss_weight=0.25, - # (float) The weight of policy loss. - policy_loss_weight=1, - # (float) The weight of ssl (self-supervised learning) loss. - ssl_loss_weight=0, - cos_lr_scheduler=False, - piecewise_decay_lr_scheduler=False, - # (bool) Whether to use piecewise constant learning rate decay. - # i.e. lr: 0.2 -> 0.02 -> 0.002 - lr_piecewise_constant_decay=False, - # (int) The number of final training iterations to control lr decay, which is only used for manually decay. - threshold_training_steps_for_final_lr=int(5e4), - # (bool) Whether to use manually decayed temperature. - manual_temperature_decay=False, - # (int) The number of final training iterations to control temperature, which is only used for manually decay. - threshold_training_steps_for_final_temperature=int(1e5), - # (float) The fixed temperature value for MCTS action selection, which is used to control the exploration. - # The larger the value, the more exploration. This value is only used when manual_temperature_decay=False. - fixed_temperature_value=0.25, - # (bool) Whether to use the true chance in MCTS in some environments with stochastic dynamics, such as 2048. - use_ture_chance_label_in_chance_encoder=False, - - # ****** Priority ****** - # (bool) Whether to use priority when sampling training data from the buffer. - use_priority=False, - # (float) The degree of prioritization to use. A value of 0 means no prioritization, - # while a value of 1 means full prioritization. - priority_prob_alpha=0.6, - # (float) The degree of correction to use. A value of 0 means no correction, - # while a value of 1 means full correction. - priority_prob_beta=0.4, - # (int) The initial Env Steps for training. - train_start_after_envsteps=int(0), - - # ****** UCB ****** - # (float) The alpha value used in the Dirichlet distribution for exploration at the root node of search tree. - root_dirichlet_alpha=0.3, - # (float) The noise weight at the root node of the search tree. - root_noise_weight=0.25, - - # ****** Explore by random collect ****** - # (int) The number of episodes to collect data randomly before training. - random_collect_episode_num=0, - - # ****** Explore by eps greedy ****** - eps=dict( - # (bool) Whether to use eps greedy exploration in collecting data. - eps_greedy_exploration_in_collect=False, - # (str) The type of decaying epsilon. Options are 'linear', 'exp'. - type='linear', - # (float) The start value of eps. - start=1., - # (float) The end value of eps. - end=0.05, - # (int) The decay steps from start to end eps. - decay=int(1e5), - ), - ) - - def default_model(self) -> Tuple[str, List[str]]: - """ - Overview: - Return this algorithm's default model setting for demonstration. - Returns: - - model_info (:obj:`Tuple[str, List[str]]`): A tuple containing the model name and a list of import paths. - - model_type (:obj:`str`): The model type used in this algorithm, registered in ModelRegistry. - - import_names (:obj:`List[str]`): The list of model class paths used in this algorithm. - .. note:: - Users can define and use customized network models, but they must adhere to the same interface definition - as indicated by the import_names path. For multi-task UniZero, this is ``lzero.model.unizero_model_multitask.UniZeroMTModel``. - """ - # NOTE: This specifies the default multi-task model. - return 'UniZeroMTModel', ['lzero.model.unizero_model_multitask'] - - def _init_learn(self) -> None: - """ - Overview: - Initializes the learn mode. This method is called by ``self.__init__``. - It sets up the learn model, optimizer, target model, and other utilities required for training. - """ - if self._cfg.optim_type == 'SGD': - # --- 改为SGD优化器 --- - self._optimizer_world_model = torch.optim.SGD( - self._model.world_model.parameters(), - lr=self._cfg.learning_rate, # 初始学习率,在配置中设为 0.2 - momentum=self._cfg.momentum, # 在配置中设为 0.9 - weight_decay=self._cfg.weight_decay # 在配置中设为 1e-4 - ) - elif self._cfg.optim_type == 'AdamW': - # NOTE: nanoGPT optimizer - self._optimizer_world_model = configure_optimizers_nanogpt( - model=self._model.world_model, - learning_rate=self._cfg.learning_rate, - weight_decay=self._cfg.weight_decay, - device_type=self._cfg.device, - betas=(0.9, 0.95), - ) - elif self._cfg.optim_type == 'AdamW_mix_lr_wdecay': - self._optimizer_world_model = configure_optimizer_unizero( - model=self._model.world_model, - learning_rate=self._cfg.learning_rate, # 使用一个合理的AdamW基础学习率 - weight_decay=self._cfg.weight_decay, - device_type=self._cfg.device, - betas=(0.9, 0.95), - ) - - if self._cfg.cos_lr_scheduler: - from torch.optim.lr_scheduler import CosineAnnealingLR - # TODO: check the total training steps - # self.lr_scheduler = CosineAnnealingLR(self._optimizer_world_model, 1e5, eta_min=0, last_epoch=-1) - total_iters = self._cfg.get('total_iterations', 500000) # 500k iter - # final_lr = self._cfg.get('final_learning_rate', 0.0) - final_lr = self._cfg.get('final_learning_rate', 1e-6) - - self.lr_scheduler = CosineAnnealingLR( - self._optimizer_world_model, - T_max=total_iters, - eta_min=final_lr - ) - print(f"CosineAnnealingLR enabled: T_max={total_iters}, eta_min={final_lr}") - - - if self._cfg.piecewise_decay_lr_scheduler: - from torch.optim.lr_scheduler import LambdaLR - max_step = self._cfg.threshold_training_steps_for_final_lr - # NOTE: the 1, 0.1, 0.01 is the decay rate, not the lr. - lr_lambda = lambda step: 1 if step < max_step * 0.5 else (0.1 if step < max_step else 0.01) # noqa - self.lr_scheduler = LambdaLR(self._optimizer_world_model, lr_lambda=lr_lambda) - - - # Use a deep copy for the target model. - self._target_model = copy.deepcopy(self._model) - # Ensure that the installed torch version is >= 2.0 for torch.compile. - assert int(''.join(filter(str.isdigit, torch.__version__))) >= 200, "We need torch version >= 2.0" - self._model = torch.compile(self._model) - self._target_model = torch.compile(self._target_model) - - # Wrap the target model for soft updates (momentum-based). - self._target_model = model_wrap( - self._target_model, - wrapper_name='target', - update_type='momentum', - update_kwargs={'theta': self._cfg.target_update_theta} - ) - self._learn_model = self._model - - if self._cfg.use_augmentation: - self.image_transforms = ImageTransforms( - self._cfg.augmentation, - image_shape=(self._cfg.model.observation_shape[1], self._cfg.model.observation_shape[2]) - ) - - self.value_support = DiscreteSupport(*self._cfg.model.value_support_range, self._cfg.device) - self.reward_support = DiscreteSupport(*self._cfg.model.reward_support_range, self._cfg.device) - self.value_inverse_scalar_transform_handle = InverseScalarTransform(self.value_support, self._cfg.model.categorical_distribution) - self.reward_inverse_scalar_transform_handle = InverseScalarTransform(self.reward_support, self._cfg.model.categorical_distribution) - - self.intermediate_losses = defaultdict(float) - self.l2_norm_before = 0. - self.l2_norm_after = 0. - self.grad_norm_before = 0. - self.grad_norm_after = 0. - - # Create a WrappedModel instance. - # This is used for gradient correction methods where gradients of shared parameters are managed. - # In this setup, all parameters are considered shared and subject to correction. - # wrapped_model = WrappedModel( - # self._learn_model.world_model, - # ) - - self.task_id = self._cfg.task_id - self.task_num_for_current_rank = self._cfg.task_num - - print(f'self._cfg.only_use_moco_stats:{self._cfg.only_use_moco_stats}') - if self._cfg.use_moco or self._cfg.only_use_moco_stats: - # The prediction heads' gradients are not corrected. - self.wrapped_model = WrappedModelV2( - # TODO: This assumes the tokenizer has an encoder attribute which is a list. This might need to be more robust. - self._learn_model.world_model.tokenizer.encoder[0], - self._learn_model.world_model.transformer, - self._learn_model.world_model.pos_emb, - self._learn_model.world_model.task_emb, - self._learn_model.world_model.act_embedding_table, - ) - - # Alternative setup: The head and tokenizer.encoder gradients are not corrected. - # wrapped_model = WrappedModelV3( - # self._learn_model.world_model.transformer, - # self._learn_model.world_model.pos_emb, - # self._learn_model.world_model.task_emb, - # self._learn_model.world_model.act_embedding_table, - # ) - - # Pass the wrapped_model as `shared_module` to the gradient correction method. - # ========= Initialize MoCo/CAGrad parameters ========= - if self._cfg.moco_version=="v0": - # This version is only compatible with single-GPU training. - self.grad_correct = GradCorrect(self.wrapped_model, self._cfg.total_task_num, self._cfg.device, self._cfg.multi_gpu) - self.grad_correct.init_param() - self.grad_correct.rep_grad = False - elif self._cfg.moco_version=="v1": - cfg_moco = MoCoCfg( - beta0=0.9, beta_sigma=0.95, - gamma0=0.1, gamma_sigma=0.95, - rho=0.01, stat_interval=10000) - self.grad_correct = FastMoCo( - shared_module=self.wrapped_model, - world_task_num=self._cfg.total_task_num, # Total number of tasks globally - device=self._cfg.device, - multi_gpu=self._cfg.multi_gpu, - cfg=cfg_moco, - ) - - # Cache for plasticity-related metrics from the previous frame. - self._prev_plasticity_metrics = dict( - dormant_ratio_encoder = 0.0, - dormant_ratio_transformer = 0.0, - dormant_ratio_head = 0.0, - avg_weight_mag_encoder = 0.0, - avg_weight_mag_transformer = 0.0, - avg_weight_mag_head = 0.0, - e_rank_last_linear = 0.0, - e_rank_sim_norm = 0.0, - ) - - # ==================== START: 目标熵正则化初始化 ==================== - # 从配置中读取是否启用自适应alpha,并提供一个默认值 - self.use_adaptive_entropy_weight = self._cfg.get('use_adaptive_entropy_weight', True) - - # 在 _init_learn 中增加配置 - self.target_entropy_start_ratio = self._cfg.get('target_entropy_start_ratio', 0.98) - self.target_entropy_end_ratio = self._cfg.get('target_entropy_end_ratio', 0.7) - self.target_entropy_decay_steps = self._cfg.get('target_entropy_decay_steps', 200000) # 例如,在200k步内完成退火 2M envsteps - - if self.use_adaptive_entropy_weight: - # 1. 设置目标熵。对于离散动作空间,一个常见的启发式设置是动作空间维度的负对数乘以一个系数。 - # 这个系数(例如0.98)可以作为一个超参数。 - action_space_size = self._cfg.model.action_space_size - self.target_entropy = -np.log(1.0 / action_space_size) * 0.98 - - # 2. 初始化一个可学习的 log_alpha 参数。 - # 初始化为0,意味着初始的 alpha = exp(0) = 1.0。 - self.log_alpha = torch.nn.Parameter(torch.zeros(1, device=self._cfg.device), requires_grad=True) - - # 3. 为 log_alpha 创建一个专属的优化器。 - # 使用与主优化器不同的、较小的学习率(例如1e-4)通常更稳定。 - alpha_lr = self._cfg.get('adaptive_entropy_alpha_lr', 1e-4) - self.alpha_optimizer = torch.optim.Adam([self.log_alpha], lr=alpha_lr) - - print("="*20) - print(">>> 目标熵正则化 (自适应Alpha) 已启用 <<<") - print(f" 目标熵 (Target Entropy): {self.target_entropy:.4f}") - print(f" Alpha 优化器学习率: {alpha_lr:.2e}") - print("="*20) - # ===================== END: 目标熵正则化初始化 ===================== - - self.latent_norm_clip_threshold = self._cfg.get('latent_norm_clip_threshold', 30.0) - # ==================== START: 初始化 Encoder-Clip Annealing 参数 ==================== - self.use_encoder_clip_annealing = self._cfg.get('use_encoder_clip_annealing', False) - if self.use_encoder_clip_annealing: - self.encoder_clip_anneal_type = self._cfg.get('encoder_clip_anneal_type', 'cosine') - self.encoder_clip_start = self._cfg.get('encoder_clip_start_value', 30.0) - self.encoder_clip_end = self._cfg.get('encoder_clip_end_value', 10.0) - self.encoder_clip_anneal_steps = self._cfg.get('encoder_clip_anneal_steps', 200000) - - print("="*20) - print(">>> Encoder-Clip 退火已启用 <<<") - print(f" 类型: {self.encoder_clip_anneal_type}") - print(f" 范围: {self.encoder_clip_start} -> {self.encoder_clip_end}") - print(f" 步数: {self.encoder_clip_anneal_steps}") - print("="*20) - else: - # 如果不启用退火,则使用固定的 clip 阈值 - self.latent_norm_clip_threshold = self._cfg.get('latent_norm_clip_threshold', 30.0) - # ===================== END: 初始化 Encoder-Clip Annealing 参数 ===================== - - # --- NEW: Policy Label Smoothing Parameters --- - self.policy_ls_eps_start = self._cfg.get('policy_ls_eps_start', 0.05) # TODO policy_label_smoothing_eps_start 越大的action space需要越大的eps - self.policy_ls_eps_end = self._cfg.get('policy_label_smoothing_eps_end ', 0.01) # TODO policy_label_smoothing_eps_start - self.policy_ls_eps_decay_steps = self._cfg.get('policy_ls_eps_decay_steps ', 50000) # TODO 50k - print(f"self.policy_ls_eps_start:{self.policy_ls_eps_start}") - - @staticmethod - def _is_zero(x: Union[float, torch.Tensor], eps: float = 1e-8) -> bool: - """ - Overview: - Checks if a scalar or a 0-D tensor can be considered zero within a small tolerance. - Arguments: - - x (:obj:`Union[float, torch.Tensor]`): The input value to check. - - eps (:obj:`float`): The tolerance for checking against zero. - Returns: - - (:obj:`bool`): True if the value is close to zero, False otherwise. - """ - if isinstance(x, torch.Tensor): - return torch.all(torch.abs(x) < eps).item() - return abs(x) < eps - - def _retain_prev_if_zero(self, name: str, - value: Union[float, torch.Tensor]) -> Union[float, torch.Tensor]: - """ - Overview: - If the current `value` is close to zero, returns the cached value from the previous frame. - Otherwise, it updates the cache with the current value and returns it. This is useful for - metrics that are computed intermittently. - Arguments: - - name (:obj:`str`): The name of the metric to cache. - - value (:obj:`Union[float, torch.Tensor]`): The current value of the metric. - Returns: - - (:obj:`Union[float, torch.Tensor]`): The retained or current value. - """ - if self._is_zero(value): - # Directly return the previous value (can be float or tensor). - return self._prev_plasticity_metrics[name] - else: - # Update the cache and return the current value. - self._prev_plasticity_metrics[name] = value - return value - - - #@profile - def _forward_learn(self, data: Tuple[torch.Tensor], task_weights=None, train_iter=None, ignore_grad=False) -> Dict[str, Union[float, int]]: - """ - Overview: - The forward function for learning in the policy. This is the core of the training process. - Data is sampled from the replay buffer, losses are calculated, and the model is updated via backpropagation. - Arguments: - - data (:obj:`Tuple[torch.Tensor]`): A tuple of data batches, where each element corresponds to a different task. - - task_weights (:obj:`Any`, optional): Optional weights for each task's loss. Not currently used. - - ignore_grad (:obj:`bool`): If True, gradients are zeroed out after computation, effectively skipping the update. - Returns: - - info_dict (:obj:`Dict[str, Union[float, int]]`): A dictionary containing current learning losses and statistics for logging. - """ - self._learn_model.train() - self._target_model.train() - - # Lists to store metrics for each task within the batch. - obs_loss_multi_task = [] - reward_loss_multi_task = [] - policy_loss_multi_task = [] - value_loss_multi_task = [] - latent_recon_loss_multi_task = [] - perceptual_loss_multi_task = [] - orig_policy_loss_multi_task = [] - policy_entropy_multi_task = [] - weighted_total_loss = 0.0 # Initialize to 0.0 to avoid in-place operations. - - latent_state_l2_norms_multi_task = [] - average_target_policy_entropy_multi_task = [] - value_priority_multi_task = [] - value_priority_mean_multi_task = [] - - # Metrics for network plasticity analysis. - dormant_ratio_encoder_multi_task = [] - dormant_ratio_transformer_multi_task = [] - dormant_ratio_head_multi_task = [] - avg_weight_mag_encoder_multi_task = [] - avg_weight_mag_transformer_multi_task = [] - avg_weight_mag_head_multi_task = [] - e_rank_last_linear_multi_task = [] - e_rank_sim_norm_multi_task = [] - - # --- NEW: Calculate current epsilon for policy --- - # if self.policy_ls_eps_start > 0: - # progress = min(1.0, train_iter / self.policy_ls_eps_decay_steps) - # current_policy_label_eps = self.policy_ls_eps_start * (1 - progress) + self.policy_ls_eps_end * progress - # else: - # current_policy_label_eps = 0.0 - current_policy_label_eps = 0.01 - - # 新增一个列表来收集当前批次中所有任务的真实全局ID - global_task_ids_in_batch = [] - alpha_loss = None - - losses_list = [] # Used to store the loss tensor for each task, required by gradient correction methods. - for task_id, data_one_task in enumerate(data): - current_batch, target_batch, task_id = data_one_task # task_id 是真实的全局ID - - # 将真实的全局ID添加到列表中 - global_task_ids_in_batch.append(task_id) - - # TODO: Adapt RoPE for multitask settings (using timestep_batch). - obs_batch_ori, action_batch, target_action_batch, mask_batch, indices, weights, make_time, timestep_batch = current_batch - target_reward, target_value, target_policy = target_batch - - # Prepare observations based on frame stack number. - if self._cfg.model.frame_stack_num == 4: - obs_batch, obs_target_batch = prepare_obs_stack_for_unizero(obs_batch_ori, self._cfg) - else: - obs_batch, obs_target_batch = prepare_obs(obs_batch_ori, self._cfg) - - # Apply augmentations if needed. - if self._cfg.use_augmentation: - obs_batch = self.image_transforms.transform(obs_batch) - if self._cfg.model.self_supervised_learning_loss: - obs_target_batch = self.image_transforms.transform(obs_target_batch) - - # Prepare action batch and convert to a torch tensor. - action_batch = torch.from_numpy(action_batch).to(self._cfg.device).unsqueeze( - -1).long() # For discrete action space. - data_list = [mask_batch, target_reward.astype('float32'), target_value.astype('float32'), target_policy, - weights] - mask_batch, target_reward, target_value, target_policy, weights = to_torch_float_tensor(data_list, - self._cfg.device) - - cur_batch_size = target_reward.size(0) # Run-time batch size. - - target_reward = target_reward.view(cur_batch_size, -1) - target_value = target_value.view(cur_batch_size, -1) - - # Transform scalar rewards and values to their scaled representations. - transformed_target_reward = scalar_transform(target_reward) - transformed_target_value = scalar_transform(target_value) - - # Convert scaled representations to categorical distributions. - # target_reward_categorical = phi_transform(self.reward_support, transformed_target_reward) - # target_value_categorical = phi_transform(self.value_support, transformed_target_value) - - target_reward_categorical = phi_transform(self.reward_support, transformed_target_reward, label_smoothing_eps= self._cfg.label_smoothing_eps) - target_value_categorical = phi_transform(self.value_support, transformed_target_value, label_smoothing_eps=self._cfg.label_smoothing_eps) - - - # Prepare the batch for the transformer-based world model. - batch_for_gpt = {} - if isinstance(self._cfg.model.observation_shape, int) or len(self._cfg.model.observation_shape) == 1: - batch_for_gpt['observations'] = torch.cat((obs_batch, obs_target_batch), dim=1).reshape( - cur_batch_size, -1, self._cfg.model.observation_shape) - elif len(self._cfg.model.observation_shape) == 3: - batch_for_gpt['observations'] = torch.cat((obs_batch, obs_target_batch), dim=1).reshape( - cur_batch_size, -1, *self._cfg.model.observation_shape) - - batch_for_gpt['actions'] = action_batch.squeeze(-1) - batch_for_gpt['rewards'] = target_reward_categorical[:, :-1] - batch_for_gpt['mask_padding'] = mask_batch == 1.0 # 0 means invalid padding data. - batch_for_gpt['mask_padding'] = batch_for_gpt['mask_padding'][:, :-1] - batch_for_gpt['observations'] = batch_for_gpt['observations'][:, :-1] - batch_for_gpt['ends'] = torch.zeros(batch_for_gpt['mask_padding'].shape, dtype=torch.long, - device=self._cfg.device) - batch_for_gpt['target_value'] = target_value_categorical[:, :-1] - batch_for_gpt['target_policy'] = target_policy[:, :-1] - batch_for_gpt['scalar_target_value'] = target_value - - # Extract valid target policy data and compute its entropy. - valid_target_policy = batch_for_gpt['target_policy'][batch_for_gpt['mask_padding']] - target_policy_entropy = -torch.sum(valid_target_policy * torch.log(valid_target_policy + 1e-9), dim=-1) - average_target_policy_entropy = target_policy_entropy.mean().item() - - # Update world model and compute losses. - intermediate_losses = defaultdict(float) - # losses = self._learn_model.world_model.compute_loss( - # batch_for_gpt, self._target_model.world_model.tokenizer, self.value_inverse_scalar_transform_handle, task_id=task_id - # ) - - losses = self._learn_model.world_model.compute_loss( - batch_for_gpt, self._target_model.world_model.tokenizer, self.value_inverse_scalar_transform_handle, current_policy_label_eps=current_policy_label_eps, task_id=task_id - ) - - # ==================== START MODIFICATION 2 ==================== - # Extract the calculated value_priority from the returned losses. - value_priority_tensor = losses.intermediate_losses['value_priority'] - # Convert to numpy array for the replay buffer, adding a small epsilon. - value_priority_np = value_priority_tensor.detach().cpu().numpy() + 1e-6 - # ===================== END MODIFICATION 2 ===================== - - - # TODO: Accumulate the weighted total loss. This assumes the loss from `compute_loss` is already weighted. - weighted_total_loss += losses.loss_total # NOTE:+= - - # TODO: Add assertions to check for NaN or Inf values in the loss if needed for debugging. - # assert not torch.isnan(losses.loss_total).any(), "Loss contains NaN values" - # assert not torch.isinf(losses.loss_total).any(), "Loss contains Inf values" - - # TODO: Append the total loss for this task, used by MoCo. - losses_list.append(losses.loss_total) - - for loss_name, loss_value in losses.intermediate_losses.items(): - intermediate_losses[f"{loss_name}"] = loss_value - - - - obs_loss = intermediate_losses['loss_obs'] - reward_loss = intermediate_losses['loss_rewards'] - policy_loss = intermediate_losses['loss_policy'] - orig_policy_loss = intermediate_losses['orig_policy_loss'] - policy_entropy = intermediate_losses['policy_entropy'] - value_loss = intermediate_losses['loss_value'] - latent_recon_loss = intermediate_losses['latent_recon_loss'] - perceptual_loss = intermediate_losses['perceptual_loss'] - latent_state_l2_norms = intermediate_losses['latent_state_l2_norms'] - - # 从 losses 对象中提取策略熵 - # ==================== START: 目标熵正则化更新逻辑 ==================== - current_alpha = self._cfg.model.world_model_cfg.policy_entropy_weight # 默认使用固定值 - if self.use_adaptive_entropy_weight: - # --- 动态计算目标熵 (这部分逻辑是正确的,予以保留) --- - progress = min(1.0, train_iter / self.target_entropy_decay_steps) - current_ratio = self.target_entropy_start_ratio * (1 - progress) + self.target_entropy_end_ratio * progress - action_space_size = self._cfg.model.action_space_size - # 注意:我们将 target_entropy 定义为正数,更符合直觉 - current_target_entropy = -np.log(1.0 / action_space_size) * current_ratio - - # --- 计算 alpha_loss (已修正符号) --- - # 这是核心修正点:去掉了最前面的负号 - # detach() 仍然是关键,确保 alpha_loss 的梯度只流向 log_alpha - alpha_loss = (self.log_alpha * (policy_entropy.detach() - current_target_entropy)).mean() # NOTE:= - - # # --- 更新 log_alpha --- - self.alpha_optimizer.zero_grad() - alpha_loss.backward() - self.alpha_optimizer.step() - # --- [优化建议] 增加 log_alpha 裁剪作为安全措施 --- - with torch.no_grad(): - # 将 alpha 限制在例如 [1e-4, 10.0] 的范围内 - self.log_alpha.clamp_(np.log(1e-4), np.log(10.0)) - - # --- 使用当前更新后的 alpha (截断梯度流) --- - current_alpha = self.log_alpha.exp().detach() - - # 重新计算加权的策略损失和总损失 - # 注意:这里的 policy_entropy 已经是一个batch的平均值 - weighted_policy_loss = orig_policy_loss - current_alpha * policy_entropy - # 重新构建总损失 (不使用 losses.loss_total) - # 确保这里的权重与 LossWithIntermediateLosses 类中的计算方式一致 - self.obs_loss_weight = 10 - self.value_loss_weight = 0.5 - self.reward_loss_weight = 1. - self.policy_loss_weight = 1. - self.ends_loss_weight = 0. - total_loss = ( - self.reward_loss_weight * reward_loss + - self.value_loss_weight * value_loss + - self.policy_loss_weight * weighted_policy_loss + - self.obs_loss_weight * obs_loss # 假设 ssl_loss_weight 是 obs_loss 的权重 - # ... 如果还有其他损失项,也加进来 ... - ) - weighted_total_loss += (weights * total_loss).mean() # NOTE:+= - # ===================== END: 目标熵正则化更新逻辑 ===================== - - # ============ For value-based priority calculation ============ - # TODO: The following section for calculating value_priority is commented out. - # If re-enabled, ensure it correctly computes L1 loss between predicted and target values - # and handles CPU/Numpy conversion properly. - # original_value = self.value_inverse_scalar_transform_handle(logits_value.reshape(-1, 101)).reshape( - # batch_for_gpt['observations'].shape[0], batch_for_gpt['observations'].shape[1], 1) - # value_priority = torch.nn.L1Loss(reduction='none')(original_value.squeeze(-1)[:,0], target_value[:, 0]) - # value_priority = value_priority.data.cpu().numpy() + 1e-6 - # value_priority = torch.tensor(0., device=self._cfg.device) - # ============ End of value priority section ============ - - # Metrics related to network plasticity. - # Use the helper function to retain the previous value if the current one is zero. - dormant_ratio_encoder = self._retain_prev_if_zero( - 'dormant_ratio_encoder', - intermediate_losses['dormant_ratio_encoder']) - dormant_ratio_transformer = self._retain_prev_if_zero( - 'dormant_ratio_transformer', - intermediate_losses['dormant_ratio_transformer']) - dormant_ratio_head = self._retain_prev_if_zero( - 'dormant_ratio_head', - intermediate_losses['dormant_ratio_head']) - avg_weight_mag_encoder = self._retain_prev_if_zero( - 'avg_weight_mag_encoder', - intermediate_losses['avg_weight_mag_encoder']) - avg_weight_mag_transformer = self._retain_prev_if_zero( - 'avg_weight_mag_transformer', - intermediate_losses['avg_weight_mag_transformer']) - avg_weight_mag_head = self._retain_prev_if_zero( - 'avg_weight_mag_head', - intermediate_losses['avg_weight_mag_head']) - e_rank_last_linear = self._retain_prev_if_zero( - 'e_rank_last_linear', - intermediate_losses['e_rank_last_linear']) - e_rank_sim_norm = self._retain_prev_if_zero( - 'e_rank_sim_norm', - intermediate_losses['e_rank_sim_norm']) - - # Append all metrics for this task to their respective lists. - obs_loss_multi_task.append(obs_loss) - reward_loss_multi_task.append(reward_loss) - policy_loss_multi_task.append(policy_loss) - orig_policy_loss_multi_task.append(orig_policy_loss) - policy_entropy_multi_task.append(policy_entropy) - value_loss_multi_task.append(value_loss) - latent_recon_loss_multi_task.append(latent_recon_loss) - perceptual_loss_multi_task.append(perceptual_loss) - latent_state_l2_norms_multi_task.append(latent_state_l2_norms) - value_priority_multi_task.append(value_priority_tensor) - value_priority_mean_multi_task.append(value_priority_tensor.mean().item()) - - # Append plasticity metrics. - dormant_ratio_encoder_multi_task.append(dormant_ratio_encoder) - dormant_ratio_transformer_multi_task.append(dormant_ratio_transformer) - dormant_ratio_head_multi_task.append(dormant_ratio_head) - avg_weight_mag_encoder_multi_task.append(avg_weight_mag_encoder) - avg_weight_mag_transformer_multi_task.append(avg_weight_mag_transformer) - avg_weight_mag_head_multi_task.append(avg_weight_mag_head) - e_rank_last_linear_multi_task.append(e_rank_last_linear) - e_rank_sim_norm_multi_task.append(e_rank_sim_norm) - - - # Core learn model update step. - self._optimizer_world_model.zero_grad() - - # Assuming losses_list is a list of tensors with gradients, e.g., [loss1, loss2, ...]. - if self._cfg.use_moco: - # Call MoCo's backward method, which handles gradient correction internally. - if self._cfg.moco_version=="v0": - lambd, stats = self.grad_correct.backward(losses=losses_list, **self._cfg.grad_correct_params) - elif self._cfg.moco_version=="v1": - lambd, stats = self.grad_correct.backward(losses_list) - - elif self._cfg.only_use_moco_stats: - # Only compute MoCo stats without applying gradient correction. - lambd, stats = self.grad_correct.backward(losses=losses_list, **self._cfg.grad_correct_params) - # Each rank performs its own backpropagation. - weighted_total_loss.backward() - else: - # If not using gradient correction, each rank performs standard backpropagation. - lambd = torch.tensor([0. for _ in range(self.task_num_for_current_rank)], device=self._cfg.device) - weighted_total_loss.backward() - - - # ----------------------------------------------------------------- - # 仍然在 torch.no_grad() 环境下执行 - # ================================================================= - with torch.no_grad(): - # 1. Encoder-Clip - # ==================== START: 动态计算当前 Clip 阈值 ==================== - current_clip_value = self.latent_norm_clip_threshold # 默认使用固定值 - if self.use_encoder_clip_annealing: - progress = min(1.0, train_iter / self.encoder_clip_anneal_steps) - - if self.encoder_clip_anneal_type == 'cosine': - # 余弦调度: 从1平滑过渡到0 - cosine_progress = 0.5 * (1.0 + np.cos(np.pi * progress)) - current_clip_value = self.encoder_clip_end + \ - (self.encoder_clip_start - self.encoder_clip_end) * cosine_progress - else: # 默认为线性调度 - current_clip_value = self.encoder_clip_start * (1 - progress) + \ - self.encoder_clip_end * progress - # ===================== END: 动态计算当前 Clip 阈值 ===================== - - # 1. Encoder-Clip (使用动态计算出的 current_clip_value) - if current_clip_value > 0 and 'obs_embeddings' in losses.intermediate_losses: - obs_embeddings = losses.intermediate_losses['obs_embeddings'] - if obs_embeddings is not None: - max_latent_norm = obs_embeddings.norm(p=2, dim=-1).max() - if max_latent_norm > current_clip_value: - scale_factor = current_clip_value / max_latent_norm.item() - # 不再频繁打印,或者可以改为每隔N步打印一次 - if train_iter % 1000 == 0: - print(f"[Encoder-Clip Annealing] Iter {train_iter}: Max latent norm {max_latent_norm.item():.2f} > {current_clip_value:.2f}. Scaling by {scale_factor:.4f}.") - scale_module_weights_vectorized(self._model.world_model.tokenizer.encoder, scale_factor) - - - # For debugging purposes. - # for name, param in self._learn_model.world_model.tokenizer.encoder.named_parameters(): - # print('name, param.mean(), param.std():', name, param.mean(), param.std()) - # if param.requires_grad: - # print(name, param.grad.norm()) - - if self._cfg.analysis_sim_norm: - del self.l2_norm_before, self.l2_norm_after, self.grad_norm_before, self.grad_norm_after - self.l2_norm_before, self.l2_norm_after, self.grad_norm_before, self.grad_norm_after = self._learn_model.encoder_hook.analyze() - self._target_model.encoder_hook.clear_data() - - total_grad_norm_before_clip_wm = torch.nn.utils.clip_grad_norm_(self._learn_model.world_model.parameters(), - self._cfg.grad_clip_value) - - if ignore_grad: - # NOTE: For cases where all tasks on a GPU are solved, `train` is still called for DDP synchronization, - # but gradients should be zeroed out to prevent updates. - self._optimizer_world_model.zero_grad() - - if self._cfg.multi_gpu: - # If not using a gradient correction method that handles it, sync gradients manually. - if not self._cfg.use_moco: - self.sync_gradients(self._learn_model) - - self._optimizer_world_model.step() - - if self._cfg.cos_lr_scheduler or self._cfg.piecewise_decay_lr_scheduler: - self.lr_scheduler.step() - - # Core target model update step. - self._target_model.update(self._learn_model.state_dict()) - - if torch.cuda.is_available(): - torch.cuda.synchronize() - current_memory_allocated = torch.cuda.memory_allocated() - max_memory_allocated = torch.cuda.max_memory_allocated() - current_memory_allocated_gb = current_memory_allocated / (1024 ** 3) - max_memory_allocated_gb = max_memory_allocated / (1024 ** 3) - else: - current_memory_allocated_gb = 0. - max_memory_allocated_gb = 0. - - # Build the dictionary of return values for logging. - return_log_dict = { - 'Current_GPU': current_memory_allocated_gb, - 'Max_GPU': max_memory_allocated_gb, - 'collect_mcts_temperature': self._collect_mcts_temperature, - 'collect_epsilon': self._collect_epsilon, - 'cur_lr_world_model': self._optimizer_world_model.param_groups[0]['lr'], - 'weighted_total_loss': weighted_total_loss.item(), - 'total_grad_norm_before_clip_wm': total_grad_norm_before_clip_wm.item(), - } - - # ==================== START: 添加新日志项 ==================== - if self.use_adaptive_entropy_weight: - return_log_dict['adaptive_alpha'] = current_alpha.item() - return_log_dict['adaptive_target_entropy_ratio'] = current_ratio - return_log_dict['alpha_loss'] = alpha_loss.item() - # ==================== START: 添加新日志项 ==================== - - # Generate task-related loss dictionaries and prefix each task-related loss with "noreduce_". - multi_task_loss_dicts = { - **generate_task_loss_dict(obs_loss_multi_task, 'noreduce_obs_loss_task{}', task_id=self.task_id), #global_task_ids=global_task_ids_in_batch), # task_id=self.task_id), - **generate_task_loss_dict(latent_recon_loss_multi_task, 'noreduce_latent_recon_loss_task{}', task_id=self.task_id), - **generate_task_loss_dict(perceptual_loss_multi_task, 'noreduce_perceptual_loss_task{}', task_id=self.task_id), - **generate_task_loss_dict(latent_state_l2_norms_multi_task, 'noreduce_latent_state_l2_norms_task{}', task_id=self.task_id), - **generate_task_loss_dict(dormant_ratio_head_multi_task, 'noreduce_dormant_ratio_head_task{}', task_id=self.task_id), - - **generate_task_loss_dict(policy_loss_multi_task, 'noreduce_policy_loss_task{}', task_id=self.task_id), - **generate_task_loss_dict(orig_policy_loss_multi_task, 'noreduce_orig_policy_loss_task{}', task_id=self.task_id), - **generate_task_loss_dict(policy_entropy_multi_task, 'noreduce_policy_entropy_task{}', task_id=self.task_id), - **generate_task_loss_dict(reward_loss_multi_task, 'noreduce_reward_loss_task{}', task_id=self.task_id), - **generate_task_loss_dict(value_loss_multi_task, 'noreduce_value_loss_task{}', task_id=self.task_id), - **generate_task_loss_dict(average_target_policy_entropy_multi_task, 'noreduce_target_policy_entropy_task{}', task_id=self.task_id), - **generate_task_loss_dict(lambd, 'noreduce_lambd_task{}', task_id=self.task_id), - **generate_task_loss_dict(value_priority_multi_task, 'noreduce_value_priority_task{}', task_id=self.task_id), - **generate_task_loss_dict(value_priority_mean_multi_task, 'noreduce_value_priority_mean_task{}', task_id=self.task_id), - } - return_log_dict.update(multi_task_loss_dicts) - - - if self._learn_model.world_model.do_analysis: - # Include plasticity metrics if analysis is enabled. - plasticity_loss_dicts = { - **generate_task_loss_dict(dormant_ratio_encoder_multi_task, 'noreduce_dormant_ratio_encoder_task{}', task_id=self.task_id), - **generate_task_loss_dict(dormant_ratio_transformer_multi_task, 'noreduce_dormant_ratio_transformer_task{}', task_id=self.task_id), - **generate_task_loss_dict(dormant_ratio_head_multi_task, 'noreduce_dormant_ratio_head_task{}', task_id=self.task_id), - **generate_task_loss_dict(avg_weight_mag_encoder_multi_task, 'noreduce_avg_weight_mag_encoder_task{}', task_id=self.task_id), - **generate_task_loss_dict(avg_weight_mag_transformer_multi_task, 'noreduce_avg_weight_mag_transformer_task{}', task_id=self.task_id), - **generate_task_loss_dict(avg_weight_mag_head_multi_task, 'noreduce_avg_weight_mag_head_task{}', task_id=self.task_id), - **generate_task_loss_dict(e_rank_last_linear_multi_task, 'noreduce_e_rank_last_linear_task{}', task_id=self.task_id), - **generate_task_loss_dict(e_rank_sim_norm_multi_task, 'noreduce_e_rank_sim_norm_task{}', task_id=self.task_id), - } - # Merge the dictionaries. - return_log_dict.update(plasticity_loss_dicts) - - # Return the final loss dictionary. - return return_log_dict - - def monitor_weights_and_grads(self, model: torch.nn.Module) -> None: - """ - Overview: - A utility function to print the mean and standard deviation of weights and their gradients for each layer in a model. - Useful for debugging training issues like exploding or vanishing gradients. - Arguments: - - model (:obj:`torch.nn.Module`): The model to monitor. - """ - for name, param in model.named_parameters(): - if param.requires_grad: - print(f"Layer: {name} | " - f"Weight mean: {param.data.mean():.4f} | " - f"Weight std: {param.data.std():.4f} | " - f"Grad mean: {param.grad.mean():.4f} | " - f"Grad std: {param.grad.std():.4f}") - - def _init_collect(self) -> None: - """ - Overview: - Initializes the collect mode. This method is called by ``self.__init__``. - It sets up the collect model and MCTS utilities for data collection. - """ - self._collect_model = self._model - - # Create a copy of the configuration for collect MCTS and set a specific number of simulations. - mcts_collect_cfg = copy.deepcopy(self._cfg) - mcts_collect_cfg.num_simulations = self._cfg.collect_num_simulations - - if self._cfg.mcts_ctree: - self._mcts_collect = MCTSCtree(mcts_collect_cfg) - else: - self._mcts_collect = MCTSPtree(mcts_collect_cfg) - - self._collect_mcts_temperature = 1. - self._collect_epsilon = 0.0 - self.collector_env_num = self._cfg.collector_env_num - if self._cfg.model.model_type == 'conv': - self.last_batch_obs = torch.zeros([self.collector_env_num, self._cfg.model.observation_shape[0], 64, 64]).to(self._cfg.device) - self.last_batch_action = [-1 for i in range(self.collector_env_num)] - elif self._cfg.model.model_type == 'mlp': - self.last_batch_obs = torch.zeros([self.collector_env_num, self._cfg.model.observation_shape]).to(self._cfg.device) - self.last_batch_action = [-1 for i in range(self.collector_env_num)] - - # TODO: The num_tasks parameter is hardcoded. It should ideally be derived from the config. - def _monitor_vars_learn(self, num_tasks: int = 2) -> List[str]: - """ - Overview: - Registers variables to be monitored during training. These variables will be logged in TensorBoard. - It dynamically creates variable names for each task if `num_tasks` is provided. - Arguments: - - num_tasks (:obj:`int`): The number of tasks being trained on the current rank. - Returns: - - monitored_vars (:obj:`List[str]`): A list of strings, where each string is the name of a variable to be logged. - """ - # Basic monitored variables that do not depend on the number of tasks. - monitored_vars = [ - 'Current_GPU', - 'Max_GPU', - 'collect_epsilon', - 'collect_mcts_temperature', - 'cur_lr_world_model', - 'weighted_total_loss', - 'total_grad_norm_before_clip_wm', - - # 'value_priority', - 'adaptive_alpha', - "adaptive_target_entropy_ratio", - 'alpha_loss', - ] - - - - # Task-specific variables to be monitored. - task_specific_vars = [ - 'noreduce_obs_loss', - 'noreduce_orig_policy_loss', - 'noreduce_policy_loss', - 'noreduce_latent_recon_loss', - 'noreduce_policy_entropy', - 'noreduce_target_policy_entropy', - 'noreduce_reward_loss', - 'noreduce_value_loss', - 'noreduce_perceptual_loss', - 'noreduce_latent_state_l2_norms', - 'noreduce_lambd', - 'noreduce_value_priority_mean', - # Metrics related to network plasticity. - 'noreduce_dormant_ratio_encoder', - 'noreduce_dormant_ratio_transformer', - 'noreduce_dormant_ratio_head', - 'noreduce_avg_weight_mag_encoder', - 'noreduce_avg_weight_mag_transformer', - 'noreduce_avg_weight_mag_head', - 'noreduce_e_rank_last_linear', - 'noreduce_e_rank_sim_norm' - ] - - # Use self.task_num_for_current_rank as the number of tasks for the current rank. - num_tasks = self.task_num_for_current_rank - # If the number of tasks is provided, extend the monitored variables list with task-specific variable names. - if num_tasks is not None: - for var in task_specific_vars: - for task_idx in range(num_tasks): - monitored_vars.append(f'{var}_task{self.task_id+task_idx}') - else: - # If num_tasks is not provided, assume a single task and use the original variable names. - monitored_vars.extend(task_specific_vars) - - return monitored_vars - - #@profile - def _forward_collect( - self, - data: torch.Tensor, - action_mask: list = None, - temperature: float = 1, - to_play: List = [-1], - epsilon: float = 0.25, - ready_env_id: np.array = None, - timestep: List = [0], - task_id: int = None, - ) -> Dict: - """ - Overview: - The forward function for collecting data. It uses the model to perform MCTS search and - selects actions via sampling to encourage exploration. - Arguments: - - data (:obj:`torch.Tensor`): The input data, i.e., the current observation. - - action_mask (:obj:`list`, optional): A list of action masks for each environment. - - temperature (:obj:`float`, optional): The temperature for MCTS action selection. - - to_play (:obj:`List`, optional): A list of player IDs for each environment. - - epsilon (:obj:`float`, optional): The probability for epsilon-greedy exploration. - - ready_env_id (:obj:`np.array`, optional): An array of IDs for environments that are ready for a new action. - - timestep (:obj:`List`, optional): The current timestep in each environment. - - task_id (:obj:`int`, optional): The ID of the task for the current environments. - Returns: - - output (:obj:`Dict`): A dictionary where keys are environment IDs and values are dictionaries - containing the selected action and other MCTS statistics. - """ - self._collect_model.eval() - - self._collect_mcts_temperature = temperature - self._collect_epsilon = epsilon - active_collect_env_num = data.shape[0] - if ready_env_id is None: - ready_env_id = np.arange(active_collect_env_num) - output = {i: None for i in ready_env_id} - - with torch.no_grad(): - network_output = self._collect_model.initial_inference(self.last_batch_obs, self.last_batch_action, data, task_id=task_id) - latent_state_roots, reward_roots, pred_values, policy_logits = mz_network_output_unpack(network_output) - - pred_values = self.value_inverse_scalar_transform_handle(pred_values).detach().cpu().numpy() - latent_state_roots = latent_state_roots.detach().cpu().numpy() - - # ========================== 核心修复 ========================== - # C++ 绑定需要一个 list,即使它在 MuZero 中代表奖励。 - reward_roots = reward_roots.detach().cpu().numpy().tolist() - # =============================================================== - - policy_logits = policy_logits.detach().cpu().numpy().tolist() - - legal_actions = [[i for i, x in enumerate(action_mask[j]) if x == 1] for j in range(active_collect_env_num)] - # The main difference between collect and eval is the addition of Dirichlet noise at the root. - noises = [ - np.random.dirichlet([self._cfg.root_dirichlet_alpha] * int(sum(action_mask[j])) - ).astype(np.float32).tolist() for j in range(active_collect_env_num) - ] - if self._cfg.mcts_ctree: - # C++ MCTS tree implementation. - roots = MCTSCtree.roots(active_collect_env_num, legal_actions) - else: - # Python MCTS tree implementation. - roots = MCTSPtree.roots(active_collect_env_num, legal_actions) - - - # # 在本文件开始,通过全局变量来控制是否处于调试状态 - # global DEBUG_ENABLED;DEBUG_ENABLED = True - # import torch.distributed as dist - # if dist.get_rank() == 0 and DEBUG_ENABLED: - # print(f"rank {dist.get_rank()} 进入调试模式,输入interact,可以键入整段的python代码调试。通过设置 DEBUG_ENABLED = False, 可以跳过调试状态") - # import ipdb; ipdb.set_trace() - # # 同步点,防止其它进程早跑 - # dist.barrier() - - roots.prepare(self._cfg.root_noise_weight, noises, reward_roots, policy_logits, to_play) - self._mcts_collect.search(roots, self._collect_model, latent_state_roots, to_play, timestep= timestep, task_id=task_id) - - roots_visit_count_distributions = roots.get_distributions() - roots_values = roots.get_values() - - batch_action = [] - for i, env_id in enumerate(ready_env_id): - distributions, value = roots_visit_count_distributions[i], roots_values[i] - - if self._cfg.eps.eps_greedy_exploration_in_collect: - # Epsilon-greedy collection strategy. - action_index_in_legal_action_set, visit_count_distribution_entropy = select_action( - distributions, temperature=self._collect_mcts_temperature, deterministic=True - ) - action = np.where(action_mask[i] == 1.0)[0][action_index_in_legal_action_set] - if np.random.rand() < self._collect_epsilon: - action = np.random.choice(legal_actions[i]) - else: - # Standard collection strategy (sampling from MCTS policy). - # NOTE: `action_index_in_legal_action_set` is the index within the set of legal actions. - action_index_in_legal_action_set, visit_count_distribution_entropy = select_action( - distributions, temperature=self._collect_mcts_temperature, deterministic=False - ) - # Convert the index back to the action in the full action space. - action = np.where(action_mask[i] == 1.0)[0][action_index_in_legal_action_set] - - # ============== TODO: This section is for visualization purposes only and should be removed for training. ============== - # It forces deterministic action selection during collection. - # action_index_in_legal_action_set, visit_count_distribution_entropy = select_action( - # distributions, temperature=self._collect_mcts_temperature, deterministic=True - # ) - # action = np.where(action_mask[i] == 1.0)[0][action_index_in_legal_action_set] - # ============== End of visualization section. ============== - - output[env_id] = { - 'action': action, - 'visit_count_distributions': distributions, - 'visit_count_distribution_entropy': visit_count_distribution_entropy, - 'searched_value': value, - 'predicted_value': pred_values[i], - 'predicted_policy_logits': policy_logits[i], - } - batch_action.append(action) - - self.last_batch_obs = data - self.last_batch_action = batch_action - - # ========= TODO: This logic is currently for the `muzero_segment_collector`. ========= - if active_collect_env_num < self.collector_env_num: - # When one environment in `collect_env` finishes early, the length of `self.last_batch_obs` is reduced. - # The transformer needs the `env_id` to retrieve from the KV cache, which is complex to manage with a dynamic batch size. - # Therefore, we reset `self.last_batch_action` for all environments to -1, forcing the transformer - # to start from scratch and avoid retrieval errors. - print('==========collect_forward============') - print(f'len(self.last_batch_obs) < self.collector_env_num, {active_collect_env_num}<{self.collector_env_num}') - self._reset_collect(reset_init_data=True, task_id=task_id) - if getattr(self._cfg, 'sample_type', '') == 'episode': - print('BUG: sample_type is episode, but len(self.last_batch_obs) < self.collector_env_num') - - return output - - def _init_eval(self) -> None: - """ - Overview: - Initializes the eval mode. This method is called by ``self.__init__``. - It sets up the eval model and MCTS utilities for evaluation. - """ - self._eval_model = self._model - - # Create a copy of the configuration for eval MCTS and set a specific number of simulations. - mcts_eval_cfg = copy.deepcopy(self._cfg) - mcts_eval_cfg.num_simulations = self._cfg.eval_num_simulations - - if self._cfg.mcts_ctree: - self._mcts_eval = MCTSCtree(mcts_eval_cfg) - else: - self._mcts_eval = MCTSPtree(mcts_eval_cfg) - - self.evaluator_env_num = self._cfg.evaluator_env_num - - if self._cfg.model.model_type == 'conv': - self.last_batch_obs = torch.zeros([self.evaluator_env_num, self._cfg.model.observation_shape[0], 64, 64]).to(self._cfg.device) - self.last_batch_action = [-1 for _ in range(self.evaluator_env_num)] - elif self._cfg.model.model_type == 'mlp': - self.last_batch_obs = torch.zeros([self.evaluator_env_num, self._cfg.model.observation_shape]).to(self._cfg.device) - self.last_batch_action = [-1 for _ in range(self.evaluator_env_num)] - - #@profile - def _forward_eval(self, data: torch.Tensor, action_mask: list, to_play: int = -1, - ready_env_id: np.array = None, timestep: List = [0], task_id: int = None) -> Dict: - """ - Overview: - The forward function for evaluating the policy. It uses the model to perform MCTS search and - selects actions deterministically (choosing the one with the highest visit count). - Arguments: - - data (:obj:`torch.Tensor`): The input data, i.e., the current observation. - - action_mask (:obj:`list`): A list of action masks for each environment. - - to_play (:obj:`int`, optional): The player ID for the current turn. - - ready_env_id (:obj:`np.array`, optional): An array of IDs for environments that are ready for a new action. - - timestep (:obj:`List`, optional): The current timestep in each environment. - - task_id (:obj:`int`, optional): The ID of the task for the current environments. - Returns: - - output (:obj:`Dict`): A dictionary where keys are environment IDs and values are dictionaries - containing the selected action and other MCTS statistics. - """ - self._eval_model.eval() - active_eval_env_num = data.shape[0] - if ready_env_id is None: - ready_env_id = np.arange(active_eval_env_num) - output = {i: None for i in ready_env_id} - with torch.no_grad(): - network_output = self._eval_model.initial_inference(self.last_batch_obs_eval, self.last_batch_action, data, task_id=task_id) - latent_state_roots, reward_roots, pred_values, policy_logits = mz_network_output_unpack(network_output) - - pred_values = self.value_inverse_scalar_transform_handle(pred_values).detach().cpu().numpy() - latent_state_roots = latent_state_roots.detach().cpu().numpy() - policy_logits = policy_logits.detach().cpu().numpy().tolist() - - # ========================== 核心修复 ========================== - # C++ 绑定需要一个 list,即使它在 MuZero 中代表奖励。 - reward_roots = reward_roots.detach().cpu().numpy().tolist() # TODO============================= - # =============================================================== - - - legal_actions = [[i for i, x in enumerate(action_mask[j]) if x == 1] for j in range(active_eval_env_num)] - if self._cfg.mcts_ctree: - # C++ MCTS tree implementation. - roots = MCTSCtree.roots(active_eval_env_num, legal_actions) - else: - # Python MCTS tree implementation. - roots = MCTSPtree.roots(active_eval_env_num, legal_actions) - - # During evaluation, no noise is added to the root policy. - roots.prepare_no_noise(reward_roots, policy_logits, to_play) - self._mcts_eval.search(roots, self._eval_model, latent_state_roots, to_play, timestep= timestep, task_id=task_id) - - roots_visit_count_distributions = roots.get_distributions() - roots_values = roots.get_values() - - batch_action = [] - - for i, env_id in enumerate(ready_env_id): - distributions, value = roots_visit_count_distributions[i], roots_values[i] - - # NOTE: `deterministic=True` means we select the action with the highest visit count (argmax) - # rather than sampling, which is standard for evaluation. - action_index_in_legal_action_set, visit_count_distribution_entropy = select_action( - distributions, temperature=1, deterministic=True - ) - # Convert the index back to the action in the full action space. - action = np.where(action_mask[i] == 1.0)[0][action_index_in_legal_action_set] - - output[env_id] = { - 'action': action, - 'visit_count_distributions': distributions, - 'visit_count_distribution_entropy': visit_count_distribution_entropy, - 'searched_value': value, - 'predicted_value': pred_values[i], - 'predicted_policy_logits': policy_logits[i], - } - batch_action.append(action) - - self.last_batch_obs_eval = data - self.last_batch_action = batch_action - - return output - - #@profile - def _reset_collect(self, env_id: int = None, current_steps: int = 0, reset_init_data: bool = True, task_id: int = None) -> None: - """ - Overview: - Resets the collection process for a specific environment or all environments. - It can clear caches and reset initial data to ensure optimal performance and prevent state leakage. - Arguments: - - env_id (:obj:`int`, optional): The ID of the environment to reset. If None, the reset applies more broadly. Defaults to None. - - current_steps (:obj:`int`, optional): The current step count in the environment, used to trigger periodic cache clearing. Defaults to 0. - - reset_init_data (:obj:`bool`, optional): If True, resets the initial observation and action buffers. Defaults to True. - - task_id (:obj:`int`, optional): The task ID, currently unused in this method. Defaults to None. - """ - if reset_init_data: - self.last_batch_obs = initialize_zeros_batch( - self._cfg.model.observation_shape, - self._cfg.collector_env_num, - self._cfg.device - ) - self.last_batch_action = [-1 for _ in range(self._cfg.collector_env_num)] - # print('Collector: last_batch_obs and last_batch_action have been reset.') - - # Return immediately if env_id is not a single integer (e.g., None or a list). - # if env_id is None or isinstance(env_id, list): - # return - - # We must handle both single int and list of ints for env_id. - if env_id is not None: - if isinstance(env_id, int): - env_ids_to_reset = [env_id] - else: # Assumes it's a list - env_ids_to_reset = env_id - - # The key condition: `current_steps` is None only on the end-of-episode reset call from the collector. - if current_steps is None: - world_model = self._collect_model.world_model - for eid in env_ids_to_reset: - # Clear the specific environment's initial inference cache. - if eid < len(world_model.past_kv_cache_init_infer_envs): - world_model.past_kv_cache_init_infer_envs[eid].clear() - - print(f'>>> [Collector] Cleared KV cache for env_id: {eid} at episode end.') - - - # Determine the clear interval based on the environment's sample type. - # clear_interval = 2000 if getattr(self._cfg, 'sample_type', '') == 'episode' else 200 - clear_interval = 2000 if getattr(self._cfg, 'sample_type', '') == 'episode' else self._cfg.game_segment_length - - # Clear caches periodically to manage memory. - # if current_steps % clear_interval == 0: - if current_steps is not None and current_steps % clear_interval == 0: - - print(f'clear_interval: {clear_interval}') - - # Clear various KV caches in the collect model's world model. - world_model = self._collect_model.world_model - for kv_cache_dict_env in world_model.past_kv_cache_init_infer_envs: - kv_cache_dict_env.clear() - world_model.past_kv_cache_recurrent_infer.clear() - world_model.keys_values_wm_list.clear() - - # Free up unused GPU memory. - torch.cuda.empty_cache() - - print(f'Collector: Caches cleared for collect_model at step {current_steps} for env {env_id}.') - - # TODO: Check if resetting the target model here is correct and necessary. - self._reset_target_model() - - #@profile - def _reset_target_model(self) -> None: - """ - Overview: - Resets the target model by clearing its internal caches. This is crucial for managing memory, - especially when using transformer-based models with KV caching. - """ - # Clear various KV caches in the target model's world model. - world_model = self._target_model.world_model - for kv_cache_dict_env in world_model.past_kv_cache_init_infer_envs: - kv_cache_dict_env.clear() - world_model.past_kv_cache_recurrent_infer.clear() - world_model.keys_values_wm_list.clear() - - # Free up unused GPU memory. - torch.cuda.empty_cache() - print('Collector: Target model past_kv_cache cleared.') - - #@profile - def _reset_eval(self, env_id: int = None, current_steps: int = 0, reset_init_data: bool = True, task_id: int = None) -> None: - """ - Overview: - Resets the evaluation process for a specific environment or all environments. - Clears caches and resets initial data to ensure clean evaluation runs. - Arguments: - - env_id (:obj:`int`, optional): The ID of the environment to reset. Defaults to None. - - current_steps (:obj:`int`, optional): The current step count, used for periodic cache clearing. Defaults to 0. - - reset_init_data (:obj:`bool`, optional): If True, resets the initial observation and action buffers. Defaults to True. - - task_id (:obj:`int`, optional): The task ID. Can be used to handle different observation shapes per task. Defaults to None. - """ - if reset_init_data: - self.last_batch_obs_eval = initialize_zeros_batch( - self._cfg.model.observation_shape, - self._cfg.evaluator_env_num, - self._cfg.device - ) - # print(f'Evaluator reset: last_batch_obs_eval shape: {self.last_batch_obs_eval.shape}') - - self.last_batch_action = [-1 for _ in range(self._cfg.evaluator_env_num)] - - - # --- BEGIN ROBUST FIX --- - # This logic handles the crucial end-of-episode cache clearing for evaluation. - # The evaluator calls `_policy.reset([env_id])` when an episode is done. - if env_id is not None: - if isinstance(env_id, int): - env_ids_to_reset = [env_id] - else: # Assumes it's a list - env_ids_to_reset = env_id - - # The key condition: `current_steps` is None only on the end-of-episode reset call from the evaluator. - if current_steps is None: - world_model = self._eval_model.world_model - for eid in env_ids_to_reset: - # Clear the specific environment's initial inference cache. - if eid < len(world_model.past_kv_cache_init_infer_envs): - world_model.past_kv_cache_init_infer_envs[eid].clear() - - print(f'>>> [Evaluator] Cleared KV cache for env_id: {eid} at episode end.') - - # The recurrent cache is global. - world_model.past_kv_cache_recurrent_infer.clear() - - if hasattr(world_model, 'keys_values_wm_list'): - world_model.keys_values_wm_list.clear() - - torch.cuda.empty_cache() - return - # --- END ROBUST FIX --- - - # Determine the clear interval. - # clear_interval = 2000 if getattr(self._cfg, 'sample_type', '') == 'episode' else 200 - clear_interval = 2000 if getattr(self._cfg, 'sample_type', '') == 'episode' else self._cfg.game_segment_length - - # Clear caches periodically. - # if current_steps % clear_interval == 0: - if current_steps is not None and current_steps % clear_interval == 0: - - print(f'clear_interval: {clear_interval}') - - # Clear various KV caches in the eval model's world model. - world_model = self._eval_model.world_model - for kv_cache_dict_env in world_model.past_kv_cache_init_infer_envs: - kv_cache_dict_env.clear() - world_model.past_kv_cache_recurrent_infer.clear() - world_model.keys_values_wm_list.clear() - - # Free up unused GPU memory. - torch.cuda.empty_cache() - - print(f'Evaluator: Caches cleared for eval_model at step {current_steps} for env {env_id}.') - - - def recompute_pos_emb_diff_and_clear_cache(self) -> None: - """ - Overview: - Clears all KV caches and precomputes positional embedding matrices in the model. - This is typically called when the maximum sequence length changes. - """ - # NOTE: This must be done for both the collect and target models. - for model in [self._collect_model, self._target_model]: - model.world_model.precompute_pos_emb_diff_kv() - model.world_model.clear_caches() - torch.cuda.empty_cache() - - def _state_dict_learn(self) -> Dict[str, Any]: - """ - Overview: - Returns the state dictionary of the learn mode. - This typically includes the model, target model, and optimizer states, - which are necessary for saving and resuming training. - Returns: - - state_dict (:obj:`Dict[str, Any]`): The state dictionary for the current learning progress. - """ - return { - 'model': self._learn_model.state_dict(), - 'target_model': self._target_model.state_dict(), - 'optimizer_world_model': self._optimizer_world_model.state_dict(), - } - - # ========== NOTE: This is the original version which loads all parameters from the state_dict. ========== - # def _load_state_dict_learn(self, state_dict: Dict[str, Any]) -> None: - # """ - # Overview: - # Loads the state_dict into the policy's learn mode. - # Arguments: - # - state_dict (:obj:`Dict[str, Any]`): The state dictionary saved from a previous training session. - # """ - # self._learn_model.load_state_dict(state_dict['model']) - # self._target_model.load_state_dict(state_dict['target_model']) - # self._optimizer_world_model.load_state_dict(state_dict['optimizer_world_model']) - - # ========== NOTE: This is a pretrain-finetune version that selectively loads parameters and freezes layers. ========== - def _load_state_dict_learn(self, state_dict: Dict[str, Any], finetune_components: List[str] = []) -> None: - """ - Overview: - Loads a state_dict for fine-tuning. It excludes multi-task specific parameters - and can freeze parts of the model (e.g., encoder, transformer) based on `finetune_components`. - Arguments: - - state_dict (:obj:`Dict[str, Any]`): The state dictionary from a pre-trained model. - - finetune_components (:obj:`List[str]`, optional): A list of component names (e.g., "encoder", "transformer") - that will remain trainable. Components not in this list will have their parameters frozen. - """ - # Example configurations for fine-tuning: - # finetune_components = [] # Loads encoder & transformer, fine-tunes only heads. - # finetune_components = ['transformer'] # Loads encoder & transformer, fine-tunes transformer & heads. - finetune_components = ["representation_network", "encoder"] # Loads encoder & transformer, fine-tunes encoder & heads. - - # Define prefixes of parameters to be excluded from loading (typically multi-task heads). - exclude_prefixes = [ - '_orig_mod.world_model.head_policy_multi_task.', - '_orig_mod.world_model.head_value_multi_task.', - '_orig_mod.world_model.head_rewards_multi_task.', - '_orig_mod.world_model.head_observations_multi_task.', - '_orig_mod.world_model.task_emb.' - ] - - # Define specific parameter keys to be excluded (for special cases like task embeddings). - exclude_keys = [ - '_orig_mod.world_model.task_emb.weight', - '_orig_mod.world_model.task_emb.bias', - ] - - def filter_state_dict(state_dict_loader: Dict[str, Any], exclude_prefixes: list, exclude_keys: list = []) -> Dict[str, Any]: - """ - Filters out parameters from a state_dict based on prefixes and specific keys. - """ - filtered = {} - for k, v in state_dict_loader.items(): - if any(k.startswith(prefix) for prefix in exclude_prefixes): - print(f"Excluding parameter: {k}") # For debugging - continue - if k in exclude_keys: - print(f"Excluding specific parameter: {k}") # For debugging - continue - filtered[k] = v - return filtered - - # Filter and load the 'model' state_dict. - if 'model' in state_dict: - model_state_dict = state_dict['model'] - filtered_model_state_dict = filter_state_dict(model_state_dict, exclude_prefixes, exclude_keys) - missing_keys, unexpected_keys = self._learn_model.load_state_dict(filtered_model_state_dict, strict=False) - if missing_keys: - print(f"Missing keys when loading _learn_model: {missing_keys}") - if unexpected_keys: - print(f"Unexpected keys when loading _learn_model: {unexpected_keys}") - else: - print("No 'model' key found in the state_dict.") - - # Filter and load the 'target_model' state_dict. - if 'target_model' in state_dict: - target_model_state_dict = state_dict['target_model'] - filtered_target_model_state_dict = filter_state_dict(target_model_state_dict, exclude_prefixes, exclude_keys) - missing_keys, unexpected_keys = self._target_model.load_state_dict(filtered_target_model_state_dict, strict=False) - if missing_keys: - print(f"Missing keys when loading _target_model: {missing_keys}") - if unexpected_keys: - print(f"Unexpected keys when loading _target_model: {unexpected_keys}") - else: - print("No 'target_model' key found in the state_dict.") - - # Handle freezing/unfreezing of parameters in _learn_model based on finetune_components. - # This assumes a naming convention where component names are present in parameter names. - for name, param in self._learn_model.named_parameters(): - # Freeze the encoder if "encoder" is not in finetune_components. - if "encoder" in name and "encoder" not in finetune_components: - param.requires_grad = False - print(f"Freezing parameter: {name}") - # Freeze the representation network if "representation_network" is not in finetune_components. - elif "representation_network" in name and "representation_network" not in finetune_components: - param.requires_grad = False - print(f"Freezing parameter: {name}") - # Freeze the transformer if "transformer" is not in finetune_components. - elif "transformer" in name and "transformer" not in finetune_components: - param.requires_grad = False - print(f"Freezing parameter: {name}") - else: - # Other parameters remain trainable by default. - print(f"Parameter remains trainable: {name}") - - # NOTE: For more complex model structures, it might be better to identify modules by their class - # rather than relying on parameter names. For example: - # for module in self._learn_model.modules(): - # if isinstance(module, EncoderModule) and "encoder" not in finetune_components: - # for param in module.parameters(): - # param.requires_grad = False - - # ========== NOTE: Another pretrain-finetune version. The main difference from the above is the freezing logic and comments. ========== - # def _load_state_dict_learn(self, state_dict: Dict[str, Any]) -> None: - # """ - # Overview: - # Loads a state_dict into the policy's learn mode, excluding multi-task related parameters. - # This is intended for fine-tuning a pre-trained model on new tasks. - # Arguments: - # - state_dict (:obj:`Dict[str, Any]`): The state dictionary from a pre-trained model. - # """ - # # Define prefixes of parameters to be excluded. - # exclude_prefixes = [ - # '_orig_mod.world_model.head_policy_multi_task.', - # '_orig_mod.world_model.head_value_multi_task.', - # '_orig_mod.world_model.head_rewards_multi_task.', - # '_orig_mod.world_model.head_observations_multi_task.', - # '_orig_mod.world_model.task_emb.' - # ] - - # # Define specific parameter keys to be excluded. - # exclude_keys = [ - # '_orig_mod.world_model.task_emb.weight', - # '_orig_mod.world_model.task_emb.bias', - # ] - - # def filter_state_dict(state_dict_loader: Dict[str, Any], exclude_prefixes: list, exclude_keys: list = []) -> Dict[str, Any]: - # """ - # Filters out parameters that should not be loaded. - # """ - # filtered = {} - # for k, v in state_dict_loader.items(): - # if any(k.startswith(prefix) for prefix in exclude_prefixes): - # print(f"Excluding parameter: {k}") - # continue - # if k in exclude_keys: - # print(f"Excluding specific parameter: {k}") - # continue - # filtered[k] = v - # return filtered - - # # Filter and load the 'model' part. - # if 'model' in state_dict: - # model_state_dict = state_dict['model'] - # filtered_model_state_dict = filter_state_dict(model_state_dict, exclude_prefixes, exclude_keys) - # missing_keys, unexpected_keys = self._learn_model.load_state_dict(filtered_model_state_dict, strict=False) - # if missing_keys: - # print(f"Missing keys when loading _learn_model: {missing_keys}") - # if unexpected_keys: - # print(f"Unexpected keys when loading _learn_model: {unexpected_keys}") - # else: - # print("No 'model' key found in the state_dict.") - - # # Filter and load the 'target_model' part. - # if 'target_model' in state_dict: - # target_model_state_dict = state_dict['target_model'] - # filtered_target_model_state_dict = filter_state_dict(target_model_state_dict, exclude_prefixes, exclude_keys) - # missing_keys, unexpected_keys = self._target_model.load_state_dict(filtered_target_model_state_dict, strict=False) - # if missing_keys: - # print(f"Missing keys when loading _target_model: {missing_keys}") - # if unexpected_keys: - # print(f"Unexpected keys when loading _target_model: {unexpected_keys}") - # else: - # print("No 'target_model' key found in the state_dict.") - - # # Do not load the optimizer's state_dict when fine-tuning, as it contains state (like momentum) - # # specific to the pre-training task, which can hinder adaptation to new tasks. - # # A fresh optimizer is usually preferred. - # # if 'optimizer_world_model' in state_dict: - # # ... \ No newline at end of file diff --git a/lzero/policy/utils.py b/lzero/policy/utils.py index 1dd85d259..e9cea7d8d 100644 --- a/lzero/policy/utils.py +++ b/lzero/policy/utils.py @@ -244,7 +244,7 @@ def configure_optimizers_nanogpt( # TODO: The following code is commented out, which is crucial for a balanced pipeline. # We do not filter out parameters with `requires_grad=False` because their `requires_grad` # attribute might be set to `True` at a later stage during training. - # param_dict = {pn: p for pn, p in param_dict.items() if p.requires_grad} + param_dict = {pn: p for pn, p in param_dict.items() if p.requires_grad} # Create optimizer parameter groups. Any parameter that is 2D or higher will be weight decayed, # otherwise no. i.e. all weight tensors in matrix multiplications and embeddings will be decayed, @@ -430,7 +430,7 @@ def prepare_obs(obs_batch_ori: np.ndarray, cfg: EasyDict, task_id = None) -> Tup """ # Convert the numpy array of original observations to a PyTorch tensor and transfer it to the specified device. # Also, ensure the tensor is of the correct floating-point type for the model. - obs_batch_ori = torch.from_numpy(obs_batch_ori).to(cfg.device) + obs_batch_ori = torch.from_numpy(obs_batch_ori).to(cfg.device).float() # Calculate the dimension size to slice based on the model configuration. # For convolutional models ('conv'), use the number of frames to stack times the number of channels. @@ -493,7 +493,7 @@ def prepare_obs_bkp(obs_batch_ori: np.ndarray, cfg: EasyDict) -> Tuple[torch.Ten ---, ---, ---, ---, ---, ---, ---, ---, --- """ # obs_batch_ori = torch.from_numpy(obs_batch_ori).to(cfg.device).float() - obs_batch_ori = torch.from_numpy(obs_batch_ori).to(cfg.device) + obs_batch_ori = torch.from_numpy(obs_batch_ori).to(cfg.device).float() # ``obs_batch`` is used in ``initial_inference()``, which is the first stacked obs at timestep t in # ``obs_batch_ori``. shape is (4, (4+5)*1, 96, 96) = (4, 9, 96, 96) obs_batch = obs_batch_ori[:, 0:cfg.model.frame_stack_num * cfg.model.image_channel, :, :] diff --git a/lzero/worker/__init__.py b/lzero/worker/__init__.py index ece5213be..93a3a9d8c 100644 --- a/lzero/worker/__init__.py +++ b/lzero/worker/__init__.py @@ -3,3 +3,4 @@ from .muzero_collector import MuZeroCollector from .muzero_segment_collector import MuZeroSegmentCollector from .muzero_evaluator import MuZeroEvaluator +from .muzero_per_level_evaluator import MuZeroPerLevelEvaluator diff --git a/lzero/worker/muzero_collector.py b/lzero/worker/muzero_collector.py index 06fa3b580..4e3d0bd86 100644 --- a/lzero/worker/muzero_collector.py +++ b/lzero/worker/muzero_collector.py @@ -340,6 +340,7 @@ def collect( # --- Initializations --- collected_episode = 0 + collected_step = 0 env_nums = self._env_num retry_waiting_time = 0.05 @@ -411,7 +412,7 @@ def collect( # Policy Forward Pass # ============================================================== policy_input = { - 'x': stack_obs_tensor, + 'data': stack_obs_tensor, 'action_mask': action_mask, 'temperature': temperature, 'to_play': to_play, @@ -535,7 +536,7 @@ def collect( # --- Episode Termination Handling --- if done: collected_episode += 1 - reward = info['eval_episode_return'] + reward = info['score'] log_info = {'reward': reward, 'time': self._env_info[env_id]['time'], 'step': self._env_info[env_id]['step']} if not collect_with_pure_policy: log_info['visit_entropy'] = visit_entropies_lst[env_id] / eps_steps_lst[env_id] if eps_steps_lst[env_id] > 0 else 0 diff --git a/lzero/worker/muzero_evaluator.py b/lzero/worker/muzero_evaluator.py index 01fabd38c..345e1d46d 100644 --- a/lzero/worker/muzero_evaluator.py +++ b/lzero/worker/muzero_evaluator.py @@ -92,14 +92,13 @@ def __init__( f'./{self._exp_name}/log/{self._instance_name}', self._instance_name ) else: - # TODO(username): Refine logger setup for UniZero multitask with DDP v2. - if tb_logger is not None: - self._logger, _ = build_logger( - f'./{self._exp_name}/log/{self._instance_name}', self._instance_name, need_tb=False - ) - self._tb_logger = tb_logger + self._logger, _ = build_logger( + f'./{self._exp_name}/log/{self._instance_name}', self._instance_name, need_tb=False + ) + self._tb_logger = tb_logger self._rank = get_rank() + self._world_size = get_world_size() print(f'rank {self._rank}, self.task_id: {self.task_id}') self.reset(policy, env) @@ -199,7 +198,7 @@ def eval( envstep: int = -1, n_episode: Optional[int] = None, return_trajectory: bool = False, - ) -> Tuple[bool, Dict[str, Any]]: + ) -> Dict[str, Any]: """ Overview: Run a full evaluation process. It will evaluate the current policy, log the results, @@ -273,8 +272,12 @@ def eval( ready_env_id = set() remain_episode = n_episode eps_steps_lst = np.zeros(env_nums) + # Hard counter independent of VectorEvalMonitor's per-env deque-fullness check; guards + # against eval hanging when episodes are unevenly distributed across envs (the per-env + # deques [n//env_num, ...] may never all reach maxlen even after n_episode finishes). + total_finishes = 0 with self._timer: - while not eval_monitor.is_finished(): + while not eval_monitor.is_finished() and total_finishes < n_episode: # Check if a timeout has occurred. if self.stop_event.is_set(): self._logger.info("[EVALUATOR]: Evaluation aborted due to timeout.") @@ -286,6 +289,9 @@ def eval( ready_env_id = ready_env_id.union(set(list(new_available_env_id)[:remain_episode])) remain_episode -= min(len(new_available_env_id), remain_episode) + if not ready_env_id: + continue + # Prepare stacked observations and other inputs for the policy. stack_obs = {env_id: game_segments[env_id].get_obs() for env_id in ready_env_id} stack_obs = list(stack_obs.values()) @@ -358,12 +364,13 @@ def eval( dones[env_id] = done if episode_timestep.done: self._policy.reset([env_id]) - reward = episode_timestep.info['eval_episode_return'] - saved_info = {'eval_episode_return': episode_timestep.info['eval_episode_return']} + reward = episode_timestep.info['score'] + saved_info = {'eval_episode_return': episode_timestep.info['score']} if 'episode_info' in episode_timestep.info: saved_info.update(episode_timestep.info['episode_info']) eval_monitor.update_info(env_id, saved_info) eval_monitor.update_reward(env_id, reward) + total_finishes += 1 self._logger.info( f"[EVALUATOR] env {env_id} finished episode, final reward: {eval_monitor.get_latest_reward(env_id)}, " f"current episode count: {eval_monitor.get_current_episode()}" @@ -406,65 +413,18 @@ def eval( duration = self._timer.value episode_return = eval_monitor.get_episode_return() + mean_episode_return = np.mean(episode_return) + if mean_episode_return >= self._max_episode_return: + if save_ckpt_fn: + save_ckpt_fn('WM_ckpt_best.pth.tar') + self._max_episode_return = mean_episode_return info = { - 'train_iter': train_iter, - 'ckpt_name': f'iteration_{train_iter}.pth.tar', - 'episode_count': n_episode, - 'envstep_count': envstep_count, 'avg_envstep_per_episode': envstep_count / n_episode if n_episode > 0 else 0, - 'evaluate_time': duration, - 'avg_envstep_per_sec': envstep_count / duration if duration > 0 else 0, - 'avg_time_per_episode': n_episode / duration if duration > 0 else 0, 'reward_mean': np.mean(episode_return), 'reward_std': np.std(episode_return), 'reward_max': np.max(episode_return), 'reward_min': np.min(episode_return), } - episode_info = eval_monitor.get_episode_info() - if episode_info is not None: - info.update(episode_info) - - print(f'rank {self._rank}, self.task_id: {self.task_id}') - self._logger.info(self._logger.get_tabulate_vars_hor(info)) - - # Log to TensorBoard and WandB. - for k, v in info.items(): - if k in ['train_iter', 'ckpt_name', 'each_reward'] or not np.isscalar(v): - continue - if self.task_id is None: - self._tb_logger.add_scalar(f'{self._instance_name}_iter/{k}', v, train_iter) - self._tb_logger.add_scalar(f'{self._instance_name}_step/{k}', v, envstep) - else: - self._tb_logger.add_scalar(f'{self._instance_name}_iter_task{self.task_id}/{k}', v, train_iter) - self._tb_logger.add_scalar(f'{self._instance_name}_step_task{self.task_id}/{k}', v, envstep) - if self.policy_config.use_wandb: - wandb.log({f'{self._instance_name}_step/{k}': v}, step=envstep) - - # Check for new best performance and save checkpoint. - mean_episode_return = np.mean(episode_return) - if mean_episode_return > self._max_episode_return: - if save_ckpt_fn: - save_ckpt_fn('ckpt_best.pth.tar') - self._max_episode_return = mean_episode_return - - # Check if the stop condition is met. - stop_flag = mean_episode_return >= self._stop_value and train_iter > 0 - if stop_flag: - self._logger.info( - f"[LightZero serial pipeline] Current episode_return: {mean_episode_return} is greater than " - f"stop_value: {self._stop_value}. The agent is considered converged." - ) - - # TODO(username): Finalize DDP synchronization for evaluation results. - # if get_world_size() > 1: - # objects = [stop_flag, episode_info] - # print(f'rank {self._rank}, self.task_id: {self.task_id}') - # print('before broadcast_object_list') - # broadcast_object_list(objects, src=0) - # print('evaluator after broadcast_object_list') - # stop_flag, episode_info = objects - - episode_info = to_item(episode_info) - if return_trajectory: - episode_info['trajectory'] = game_segments - return stop_flag, episode_info \ No newline at end of file + if mean_episode_return >= self._stop_value: + stop_flag = True + return stop_flag, info \ No newline at end of file diff --git a/lzero/worker/muzero_per_level_evaluator.py b/lzero/worker/muzero_per_level_evaluator.py new file mode 100644 index 000000000..bfc9f7f1e --- /dev/null +++ b/lzero/worker/muzero_per_level_evaluator.py @@ -0,0 +1,242 @@ +import time +from collections import defaultdict +from typing import Optional, Callable, Dict, Any + +import numpy as np +import torch +from ding.torch_utils import to_ndarray, to_tensor +from ding.utils import get_rank +from ding.worker.collector.base_serial_evaluator import VectorEvalMonitor + +from lzero.mcts.buffer.game_segment import GameSegment +from lzero.mcts.utils import prepare_observation +from lzero.worker.muzero_evaluator import MuZeroEvaluator + + +class MuZeroPerLevelEvaluator(MuZeroEvaluator): + """MuZeroEvaluator with per-level TensorBoard logging. + + Tracks `level_id` from episode info and logs per-level + aggregated + reward metrics to TensorBoard with tags matching PriorZero exactly, + enabling cross-method comparison on the same TB dashboard. + """ + + def _log_per_level_tb(self, per_level_results: dict, tag_prefix: str, global_step: int) -> None: + if not per_level_results or self._tb_logger is None: + return + all_level_means = [] + for level_id in sorted(per_level_results.keys()): + rewards = per_level_results[level_id] + mean_r = np.mean(rewards) + self._tb_logger.add_scalar(f'{tag_prefix}/level_{level_id}_reward', mean_r, global_step) + all_level_means.append(mean_r) + self._tb_logger.add_scalar(f'{tag_prefix}/level_mean', np.mean(all_level_means), global_step) + self._tb_logger.add_scalar(f'{tag_prefix}/level_std', np.std(all_level_means), global_step) + self._tb_logger.add_scalar(f'{tag_prefix}/level_min', np.min(all_level_means), global_step) + self._tb_logger.add_scalar(f'{tag_prefix}/level_max', np.max(all_level_means), global_step) + + def _log_agg_tb(self, info: dict, tag_prefix: str, global_step: int) -> None: + if self._tb_logger is None: + return + for k in ['avg_envstep_per_episode', 'reward_mean', 'reward_std', 'reward_max', 'reward_min']: + if k in info: + self._tb_logger.add_scalar(f'{tag_prefix}/{k}', info[k], global_step) + + def eval( + self, + save_ckpt_fn: Optional[Callable] = None, + train_iter: int = -1, + envstep: int = -1, + n_episode: Optional[int] = None, + return_trajectory: bool = False, + ) -> Dict[str, Any]: + if torch.cuda.is_available(): + torch.cuda.set_device(get_rank()) + + episode_info = None + stop_flag = False + per_level_results = defaultdict(list) + + if get_rank() >= 0: + if n_episode is None: + n_episode = self._default_n_episode + assert n_episode is not None + envstep_count = 0 + eval_monitor = VectorEvalMonitor(self._env.env_num, n_episode) + env_nums = self._env.env_num + + self._env.reset() + self._policy.reset(task_id=self.task_id) + + init_obs = self._env.ready_obs + retry_waiting_time = 0.001 + while len(init_obs.keys()) != self._env_num: + time.sleep(retry_waiting_time) + init_obs = self._env.ready_obs + + action_mask_dict = {i: to_ndarray(init_obs[i]['action_mask']) for i in range(env_nums)} + to_play_dict = {i: to_ndarray(init_obs[i]['to_play']) for i in range(env_nums)} + timestep_dict = {} + for i in range(env_nums): + timestep_dict[i] = to_ndarray(init_obs[i].get('timestep', -1)) + + dones = np.array([False for _ in range(env_nums)]) + game_segments = [ + GameSegment( + self._env.action_space, + game_segment_length=self.policy_config.game_segment_length, + config=self.policy_config, + task_id=self.task_id, + ) for _ in range(env_nums) + ] + for i in range(env_nums): + game_segments[i].reset( + [to_ndarray(init_obs[i]['observation']) for _ in range(self.policy_config.model.frame_stack_num)] + ) + + ready_env_id = set() + remain_episode = n_episode + eps_steps_lst = np.zeros(env_nums) + total_finishes = 0 + with self._timer: + while not eval_monitor.is_finished() and total_finishes < n_episode: + if self.stop_event.is_set(): + self._logger.info("[EVALUATOR]: Evaluation aborted due to timeout.") + break + + obs = self._env.ready_obs + new_available_env_id = set(obs.keys()).difference(ready_env_id) + ready_env_id = ready_env_id.union(set(list(new_available_env_id)[:remain_episode])) + remain_episode -= min(len(new_available_env_id), remain_episode) + + if not ready_env_id: + continue + + stack_obs = {env_id: game_segments[env_id].get_obs() for env_id in ready_env_id} + stack_obs = list(stack_obs.values()) + action_mask = [action_mask_dict[env_id] for env_id in ready_env_id] + to_play = [to_play_dict[env_id] for env_id in ready_env_id] + timestep = [timestep_dict[env_id] for env_id in ready_env_id] + + stack_obs = to_ndarray(stack_obs) + stack_obs = prepare_observation(stack_obs, self.policy_config.model.model_type) + stack_obs = torch.from_numpy(stack_obs).to(self.policy_config.device).float() + + if self.task_id is None: + policy_output = self._policy.forward(stack_obs, action_mask, to_play, ready_env_id=ready_env_id, timestep=timestep) + else: + policy_output = self._policy.forward(stack_obs, action_mask, to_play, ready_env_id=ready_env_id, timestep=timestep, task_id=self.task_id) + + actions_with_env_id = {k: v['action'] for k, v in policy_output.items()} + distributions_dict_with_env_id = {k: v['visit_count_distributions'] for k, v in policy_output.items()} + if self.policy_config.sampled_algo: + root_sampled_actions_dict_with_env_id = {k: v['root_sampled_actions'] for k, v in policy_output.items()} + value_dict_with_env_id = {k: v['searched_value'] for k, v in policy_output.items()} + pred_value_dict_with_env_id = {k: v['predicted_value'] for k, v in policy_output.items()} + timestep_dict_with_env_id = {k: v.get('timestep', -1) for k, v in policy_output.items()} + visit_entropy_dict_with_env_id = {k: v['visit_count_distribution_entropy'] for k, v in policy_output.items()} + + actions, distributions_dict, value_dict, pred_value_dict, timestep_dict, visit_entropy_dict = {}, {}, {}, {}, {}, {} + if self.policy_config.sampled_algo: + root_sampled_actions_dict = {} + + for index, env_id in enumerate(ready_env_id): + actions[env_id] = actions_with_env_id.pop(env_id) + distributions_dict[env_id] = distributions_dict_with_env_id.pop(env_id) + if self.policy_config.sampled_algo: + root_sampled_actions_dict[env_id] = root_sampled_actions_dict_with_env_id.pop(env_id) + value_dict[env_id] = value_dict_with_env_id.pop(env_id) + pred_value_dict[env_id] = pred_value_dict_with_env_id.pop(env_id) + timestep_dict[env_id] = timestep_dict_with_env_id.pop(env_id) + visit_entropy_dict[env_id] = visit_entropy_dict_with_env_id.pop(env_id) + timesteps = self._env.step(actions) + timesteps = to_tensor(timesteps, dtype=torch.float32) + for env_id, episode_timestep in timesteps.items(): + obs_t, reward, done, info = episode_timestep.obs, episode_timestep.reward, episode_timestep.done, episode_timestep.info + + eps_steps_lst[env_id] += 1 + if self._policy.get_attribute('cfg').type in ['unizero', 'sampled_unizero']: + self._policy.reset(env_id=env_id, current_steps=eps_steps_lst[env_id], reset_init_data=False, task_id=self.task_id) + + game_segments[env_id].append( + actions[env_id], to_ndarray(obs_t['observation']), reward, + action_mask_dict[env_id], to_play_dict[env_id], timestep_dict[env_id], + ) + + action_mask_dict[env_id] = to_ndarray(obs_t['action_mask']) + to_play_dict[env_id] = to_ndarray(obs_t['to_play']) + timestep_dict[env_id] = to_ndarray(obs_t.get('timestep', -1)) + + dones[env_id] = done + if episode_timestep.done: + self._policy.reset([env_id]) + reward = episode_timestep.info['score'] + saved_info = {'eval_episode_return': episode_timestep.info['score']} + if 'episode_info' in episode_timestep.info: + saved_info.update(episode_timestep.info['episode_info']) + eval_monitor.update_info(env_id, saved_info) + eval_monitor.update_reward(env_id, reward) + total_finishes += 1 + + level_id = episode_timestep.info.get('level_id', None) + if level_id is not None: + per_level_results[int(level_id)].append(float(reward)) + + self._logger.info( + f"[EVALUATOR] env {env_id} finished episode (level {level_id}), " + f"reward: {reward}, count: {total_finishes}/{n_episode}" + ) + if n_episode > self._env_num: + init_obs = self._env.ready_obs + while len(init_obs.keys()) != self._env_num: + time.sleep(retry_waiting_time) + init_obs = self._env.ready_obs + + new_available_env_id = set(init_obs.keys()).difference(ready_env_id) + ready_env_id = ready_env_id.union(set(list(new_available_env_id)[:remain_episode])) + remain_episode -= min(len(new_available_env_id), remain_episode) + + action_mask_dict[env_id] = to_ndarray(init_obs[env_id]['action_mask']) + to_play_dict[env_id] = to_ndarray(init_obs[env_id]['to_play']) + timestep_dict[env_id] = to_ndarray(init_obs[env_id].get('timestep', -1)) + + game_segments[env_id] = GameSegment( + self._env.action_space, + game_segment_length=self.policy_config.game_segment_length, + config=self.policy_config, + task_id=self.task_id, + ) + game_segments[env_id].reset( + [init_obs[env_id]['observation'] for _ in range(self.policy_config.model.frame_stack_num)] + ) + + eps_steps_lst[env_id] = 0 + self._policy.reset([env_id]) + ready_env_id.remove(env_id) + + envstep_count += 1 + + episode_return = eval_monitor.get_episode_return() + mean_episode_return = np.mean(episode_return) + if mean_episode_return >= self._max_episode_return: + if save_ckpt_fn: + save_ckpt_fn('WM_ckpt_best.pth.tar') + self._max_episode_return = mean_episode_return + info = { + 'avg_envstep_per_episode': envstep_count / n_episode if n_episode > 0 else 0, + 'reward_mean': np.mean(episode_return), + 'reward_std': np.std(episode_return), + 'reward_max': np.max(episode_return), + 'reward_min': np.min(episode_return), + } + + self._log_agg_tb(info, 'eval/wm_only/agg', envstep) + self._log_per_level_tb(dict(per_level_results), 'eval/wm_only/per_level', envstep) + + self._log_agg_tb(info, 'eval/wm_only/agg_iter', train_iter) + self._log_per_level_tb(dict(per_level_results), 'eval/wm_only/per_level_iter', train_iter) + + self._log_agg_tb(info, 'deprecated/eval/wm_mcts/agg_wm_iter', train_iter) + self._log_per_level_tb(dict(per_level_results), 'deprecated/eval/wm_mcts/per_level_wm_iter', train_iter) + + return info diff --git a/lzero/worker/muzero_segment_collector.py b/lzero/worker/muzero_segment_collector.py index 7c265630b..39b154774 100644 --- a/lzero/worker/muzero_segment_collector.py +++ b/lzero/worker/muzero_segment_collector.py @@ -477,16 +477,6 @@ def collect( if self.policy_config.use_ture_chance_label_in_chance_encoder: append_kwargs['chance'] = self.chance_dict_tmp[env_id] - # [PRIORZERO-NEW] Add raw_obs_text if available in obs (not info!) - # Jericho env puts raw_obs_text in the obs dictionary - if env_id == 0 and collected_step < 5: # Debug first few steps - print(f"[OBS_DEBUG] Step {collected_step} env {env_id}: obs keys = {list(obs.keys())}") - print(f"[OBS_DEBUG] obs type = {type(obs)}") - if 'raw_obs_text' in obs: - print(f"[OBS_DEBUG] Found raw_obs_text: {str(obs['raw_obs_text'])[:100]}...") - else: - print(f"[OBS_DEBUG] NO raw_obs_text in obs!") - if 'raw_obs_text' in obs: append_kwargs['raw_obs_text'] = obs['raw_obs_text'] elif 'raw_obs_text' in info: @@ -566,7 +556,7 @@ def collect( self._total_episode_count += 1 info = { - 'reward': episode_timestep.info['eval_episode_return'], + 'reward': episode_timestep.info['score'], 'time': self._env_info[env_id]['time'], 'step': self._env_info[env_id]['step'], } diff --git a/zoo/atari/envs/atari_lightzero_env.py b/zoo/atari/envs/atari_lightzero_env.py index d40f35033..0e29f3278 100644 --- a/zoo/atari/envs/atari_lightzero_env.py +++ b/zoo/atari/envs/atari_lightzero_env.py @@ -177,7 +177,8 @@ def step(self, action: int) -> BaseEnvTimestep: self.reward = np.array(reward).astype(np.float32) self._eval_episode_return += self.reward self._timestep += 1 - if self._timestep%200==0: + # if self._timestep%200==0: + if self._timestep%20==0: logging.info(f'self._timestep: {self._timestep}') observation = self.observe() if done: diff --git a/zoo/babyai/__init__.py b/zoo/babyai/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/zoo/babyai/configs/babyai_unizero_segment_config.py b/zoo/babyai/configs/babyai_unizero_segment_config.py new file mode 100644 index 000000000..ff78844f1 --- /dev/null +++ b/zoo/babyai/configs/babyai_unizero_segment_config.py @@ -0,0 +1,225 @@ +""" +BabyAI UniZero Baseline Config (Ablation) +========================================== +Pure UniZero world-model baseline for BabyAI multi-task (18 levels). +No LLM module, no llm-prior, no vLLM — only the world model + MCTS. + +Corresponding LLM-prior experiment config: + zoo/babyai/priorzero/src/priorzero_config.py (get_priorzero_config) + +All world-model hyperparameters (embed_dim, num_layers, num_heads, batch_size, +learning_rate, replay_buffer_size, num_simulations, game_segment_length, etc.) +are kept identical to the PriorZero config for a fair ablation comparison. + +Entry point: + lzero.entry.train_unizero_segment +""" +import sys +import os +import argparse +from pathlib import Path + +sys.path.insert(0, str(Path(__file__).resolve().parents[3])) + +from easydict import EasyDict + + +def main( + env_id: str = 'babyai', + seed: int = 0, + env_addr: str = 'http://127.0.0.1:8000', + use_high_level_actions: bool = True, + max_env_step: int = int(5e5), +) -> None: + + # === Environment (aligned with PriorZero config) === + action_space_size = 20 + max_steps = 20 + wm_encoder_option = 'legacy' + wm_model_name = '/mnt/shared-storage-user/puyuan/xiongjyu/models/bge-base-en-v1.5' + + _SCALING_INTER_RL_LEVELS = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 19, 20, 21, 30, 31, 33, 36] + train_data_idx_list = [lvl - 1 for lvl in _SCALING_INTER_RL_LEVELS] + eval_data_idx_list = [lvl - 1 for lvl in _SCALING_INTER_RL_LEVELS] + + # === Collector / Evaluator (aligned with PriorZero config) === + collector_env_num = 1 + evaluator_env_num = 4 + n_episode = collector_env_num + n_evaluator_episode = len(eval_data_idx_list) # 18 + + # === World Model (aligned with PriorZero config) === + num_unroll_steps = 10 + infer_context_length = 4 + game_segment_length = 50 + num_layers = 2 + embed_dim = 768 + replay_ratio = 0.1 + batch_size = 64 + # collect_num_simulations = 50 + collect_num_simulations = 25 + eval_num_simulations = 50 + + num_simulations = 50 + # replay_buffer_size = int(3e5) + replay_buffer_size = int(5e5) + + + # ------------------------------------------------------------------ + babyai_unizero_config = dict( + env=dict( + stop_value=int(1e6), + max_steps=max_steps, + observation_shape=512, + env_id=env_id, + env_addr=env_addr, + train_data_idx_list=train_data_idx_list, + eval_data_idx_list=eval_data_idx_list, + use_high_level_actions=use_high_level_actions, + for_unizero=True, + tokenizer_path=wm_model_name, + max_action_num=action_space_size, + max_seq_len=512, + collector_env_num=collector_env_num, + evaluator_env_num=evaluator_env_num, + n_evaluator_episode=n_evaluator_episode, + manager=dict(shared_memory=False), + ), + policy=dict( + multi_gpu=False, + use_wandb=False, + learn=dict( + learner=dict( + hook=dict(save_ckpt_after_iter=1000000), + ), + ), + model=dict( + observation_shape=512, + action_space_size=action_space_size, + encoder_option=wm_encoder_option, + encoder_url=wm_model_name, + model_type="mlp", + continuous_action_space=False, + norm_type="LN", + world_model_cfg=dict( + norm_type="LN", + final_norm_option_in_head="LayerNorm", + final_norm_option_in_encoder="LayerNorm", + predict_latent_loss_type='mse', + policy_entropy_weight=5e-2, + continuous_action_space=False, + max_blocks=num_unroll_steps, + max_tokens=2 * num_unroll_steps, + context_length=2 * infer_context_length, + device="cuda", + collect_num_simulations=collect_num_simulations, + eval_num_simulations=eval_num_simulations, + action_space_size=action_space_size, + num_layers=num_layers, + num_heads=24, + embed_dim=embed_dim, + obs_type="text", + env_num=max(collector_env_num, evaluator_env_num), + decode_loss_mode=None, + latent_recon_loss_weight=0, + task_embed_option=None, + moe_in_transformer=False, + multiplication_moe_in_transformer=False, + game_segment_length=game_segment_length, + ), + ), + update_per_collect=None, + num_segments=collector_env_num, + action_type="varied_action_space", + model_path=None, + num_unroll_steps=num_unroll_steps, + reanalyze_ratio=0, + replay_ratio=replay_ratio, + batch_size=batch_size, + learning_rate=3e-4, + weight_decay=1e-4, + cos_lr_scheduler=False, + fixed_temperature_value=0.25, + manual_temperature_decay=False, + n_episode=n_episode, + train_start_after_envsteps=0, + replay_buffer_size=replay_buffer_size, + eval_freq=int(500), + eval_per_level=True, + # eval_freq=int(5e3), + collector_env_num=collector_env_num, + evaluator_env_num=evaluator_env_num, + buffer_reanalyze_freq=1 / 1000000, + reanalyze_batch_size=160, + reanalyze_partition=0.75, + device='cuda', + num_simulations=num_simulations, + game_segment_length=game_segment_length, + off_policy_degree=0, + enable_async_eval=False, + optim_type='AdamW', + grad_clip_value=10.0, + value_loss_weight=0.25, + policy_loss_weight=1.0, + reward_loss_weight=1.0, + use_adaptive_entropy_weight=False, + adaptive_entropy_alpha_lr=1e-4, + use_encoder_clip_annealing=False, + encoder_clip_anneal_type='cosine', + encoder_clip_start_value=30.0, + encoder_clip_end_value=10.0, + encoder_clip_anneal_steps=100000, + use_priority=False, + priority_prob_alpha=0.6, + priority_prob_beta=0.4, + ), + ) + babyai_unizero_config = EasyDict(babyai_unizero_config) + + babyai_unizero_create_config = dict( + env=dict( + type="babyai", + import_names=["zoo.babyai.priorzero.envs.babyai_env"], + ), + env_manager=dict(type="base"), + policy=dict( + type="unizero", + import_names=["lzero.policy.unizero"], + ), + ) + babyai_unizero_create_config = EasyDict(babyai_unizero_create_config) + + main_config = babyai_unizero_config + create_config = babyai_unizero_create_config + + main_config.exp_name = ( + f"data_unizero/babyai/babyai_unizero_18levels_" + f"nlayer{num_layers}_edim{embed_dim}_gsl{game_segment_length}_" + f"rr{replay_ratio}_bs{batch_size}_sim{num_simulations}_seed{seed}" + ) + + from lzero.entry import train_unizero_segment + + train_unizero_segment( + [main_config, create_config], + seed=seed, + model_path=main_config.policy.model_path, + max_env_step=max_env_step, + ) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description="BabyAI UniZero Baseline (no LLM)") + parser.add_argument('--seed', type=int, default=0) + parser.add_argument('--env_addr', type=str, default='http://127.0.0.1:8000') + parser.add_argument('--use_low_level_actions', action='store_true', default=False) + parser.add_argument('--max_env_step', type=int, default=int(5e5)) + args = parser.parse_args() + + os.environ['TOKENIZERS_PARALLELISM'] = 'false' + main( + seed=args.seed, + env_addr=args.env_addr, + use_high_level_actions=not args.use_low_level_actions, + max_env_step=args.max_env_step, + ) diff --git a/zoo/babyai/priorzero/README.md b/zoo/babyai/priorzero/README.md new file mode 100644 index 000000000..c49b98d8d --- /dev/null +++ b/zoo/babyai/priorzero/README.md @@ -0,0 +1,51 @@ +# PriorZero on BabyAI (AgentGym-RL) + +BabyAI is a 2D grid world with natural language missions ("go to the red ball", "pick up the blue key") and text-based observations describing the agent's local view. The AgentGym server provides **high-level semantic actions** (e.g., "go to red ball 1", "toggle and go through green closed door 1") that abstract over low-level movement, making the dynamic action space similar to Jericho. + +## Prerequisites + +Start the AgentGym BabyAI server before training: + +```bash +cd /path/to/AgentGym-RL/AgentGym/agentenv-babyai +pip install -e . +python3 -m uvicorn agentenv_babyai:app --host 0.0.0.0 --port 8000 +``` + +Verify: `curl http://127.0.0.1:8000/` should return 200. + +## Key Differences from Jericho PriorZero + +| Aspect | Jericho | BabyAI | +|---|---|---| +| Connection | Local Python `env.step()` | HTTP client → AgentGym server | +| Action space | Dynamic text commands (10-100+) | Dynamic high-level actions (3-15) or 7 atomic | +| Observation | Game engine text | Natural language grid description | +| Mission | Implicit in game context | Explicit "Your goal: ..." string | +| Reward | Sparse integer score | Continuous [0,1]: `1 - 0.9*(steps/max_steps)` | +| `data_idx` encoding | N/A (game file path) | `level = idx % 40 + 1`, `seed = idx // 40` | + +## Quick Start + +Debug mode (1 GPU, 20 steps): +```bash +cd zoo/babyai/priorzero +torchrun --nproc_per_node=1 ./src/priorzero_entry_sync_ddp.py \ + --quick_test --env_addr http://127.0.0.1:8000 --data_idx 0 +``` + +Full training (4 GPUs): +```bash +cd zoo/babyai/priorzero +bash scripts/run_priorzero_ddp.sh +``` + +Use `--use_low_level_actions` to switch to 7 atomic actions (turn left/right, move forward, pickup, drop, toggle, check). + +## Known Issues + +1. Observation "left/right" is relative to agent heading, not map coordinates +2. `pickup`/`toggle` only affect the cell directly ahead — wrong calls waste a step +3. Compound missions (PutNext, Sequence) may need mission decomposition not covered by default prompts +4. Early episodes have low reward due to step-count decay — watch the trend, not absolute values +5. If many episodes fail at reset, check the AgentGym server first (connection issues), not the algorithm diff --git a/zoo/babyai/priorzero/__init__.py b/zoo/babyai/priorzero/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/zoo/babyai/priorzero/envs/__init__.py b/zoo/babyai/priorzero/envs/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/zoo/babyai/priorzero/envs/babyai_env.py b/zoo/babyai/priorzero/envs/babyai_env.py new file mode 100644 index 000000000..c077788fe --- /dev/null +++ b/zoo/babyai/priorzero/envs/babyai_env.py @@ -0,0 +1,399 @@ +import copy +import json +import logging +import random as _random +import re +import threading +import time +from collections import OrderedDict +from typing import Any, Dict, List, Optional, Union + +import gym +import numpy as np +import torch +import requests +from requests.adapters import HTTPAdapter +from urllib3.util.retry import Retry +from transformers import AutoTokenizer + +from ding.utils import ENV_REGISTRY, get_rank, get_world_size +from ding.envs import BaseEnv, BaseEnvTimestep + + +ATOMIC_ACTIONS = [ + "turn left", "turn right", "move forward", + "pickup", "drop", "toggle", "check available actions", +] + + +class BabyAIHttpClient: + """HTTP client for AgentGym BabyAI server with retry and timeout.""" + + def __init__(self, env_addr: str, timeout: float = 10.0, max_retries: int = 3): + self._addr = env_addr.rstrip('/') + self._timeout = timeout + self._session = requests.Session() + retries = Retry( + total=max_retries, + backoff_factor=0.5, + status_forcelist=[500, 502, 503, 504], + ) + self._session.mount('http://', HTTPAdapter(max_retries=retries)) + self._session.mount('https://', HTTPAdapter(max_retries=retries)) + + def health_check(self) -> bool: + try: + r = self._session.get(f"{self._addr}/", timeout=self._timeout) + return r.status_code == 200 + except Exception: + return False + + def create(self) -> int: + r = self._session.post(f"{self._addr}/create", timeout=self._timeout) + r.raise_for_status() + data = r.json() + if "error" in data: + raise RuntimeError(f"BabyAI create error: {data['error']}") + return data["id"] + + def reset(self, env_id: int, data_idx: int) -> dict: + r = self._session.post( + f"{self._addr}/reset", + json={"id": env_id, "data_idx": data_idx}, + timeout=self._timeout, + ) + r.raise_for_status() + data = r.json() + if "error" in data: + raise RuntimeError(f"BabyAI reset error: {data['error']}") + return data + + def step(self, env_id: int, action: str) -> dict: + r = self._session.post( + f"{self._addr}/step", + json={"id": env_id, "action": action}, + timeout=self._timeout, + ) + r.raise_for_status() + data = r.json() + if "error" in data: + raise RuntimeError(f"BabyAI step error: {data['error']}") + return data + + def close(self, env_id: int): + try: + self._session.post( + f"{self._addr}/close", + json={"id": env_id}, + timeout=self._timeout, + ) + except Exception: + pass + + def close_session(self): + self._session.close() + + +def _parse_mission(obs_text: str) -> str: + """Extract mission from observation text. Format: 'Your goal: \n...'""" + if obs_text.startswith("Your goal: "): + first_line_end = obs_text.find('\n') + if first_line_end == -1: + return obs_text[len("Your goal: "):] + return obs_text[len("Your goal: "):first_line_end].strip() + return "" + + +def _parse_available_actions(obs_text: str) -> List[str]: + """Extract available actions list from observation text. + Format: '...\\nAvailable actions: ["action1", "action2", ...]' + """ + marker = "\nAvailable actions: [" + idx = obs_text.rfind(marker) + if idx == -1: + return [] + actions_str = obs_text[idx + len(marker) - 1:] # include the '[' + try: + actions = json.loads(actions_str) + if isinstance(actions, list): + return [str(a) for a in actions] + except json.JSONDecodeError: + pass + # Fallback: regex extraction + matches = re.findall(r'"([^"]*)"', actions_str) + return matches if matches else [] + + +def _strip_actions_suffix(obs_text: str) -> str: + """Remove the 'Available actions: [...]' suffix from observation text.""" + marker = "\nAvailable actions: [" + idx = obs_text.rfind(marker) + if idx != -1: + return obs_text[:idx].strip() + return obs_text + + +@ENV_REGISTRY.register('babyai') +class BabyAIEnv(BaseEnv): + """ + BabyAI environment wrapper for PriorZero. + Communicates with AgentGym BabyAI HTTP server. + Interface contract matches JerichoEnv for algorithm-layer compatibility. + """ + tokenizer: Optional[AutoTokenizer] = None + + # aligned with ScalingInter-RL: class-level counter for evaluator task cycling + _eval_cycle_counter = 0 + _eval_cycle_lock = threading.Lock() + + DEFAULT_CONFIG: Dict[str, Any] = { + 'env_addr': 'http://127.0.0.1:8000', + 'data_idx': 0, + 'data_idx_list': None, + 'train_data_idx_list': None, + 'eval_data_idx_list': None, + 'max_steps': 64, + 'max_action_num': 20, + 'tokenizer_path': 'BAAI/bge-base-en-v1.5', + 'max_seq_len': 512, + 'for_unizero': True, + 'save_replay': False, + 'use_high_level_actions': True, + 'is_collect': True, + 'collector_env_num': 1, + 'evaluator_env_num': 1, + } + + def __init__(self, cfg: Dict[str, Any]) -> None: + merged_cfg = copy.deepcopy(self.DEFAULT_CONFIG) + merged_cfg.update(cfg) + self.cfg = merged_cfg + + self.env_addr: str = self.cfg['env_addr'] + self.data_idx: int = self.cfg.get('data_idx', 0) + self.data_idx_list: Optional[List[int]] = self.cfg.get('data_idx_list', None) + self._is_collect: bool = self.cfg.get('is_collect', True) + self.max_steps: int = self.cfg['max_steps'] + self.max_action_num: int = self.cfg['max_action_num'] + self.max_seq_len: int = self.cfg['max_seq_len'] + self.for_unizero: bool = self.cfg['for_unizero'] + self.save_replay: bool = self.cfg['save_replay'] + self.use_high_level_actions: bool = self.cfg['use_high_level_actions'] + + self.world_size: int = get_world_size() + self.rank: int = get_rank() + + if BabyAIEnv.tokenizer is None: + if self.rank == 0: + BabyAIEnv.tokenizer = AutoTokenizer.from_pretrained(self.cfg['tokenizer_path']) + if self.world_size > 1: + torch.distributed.barrier() + if self.rank != 0: + BabyAIEnv.tokenizer = AutoTokenizer.from_pretrained(self.cfg['tokenizer_path']) + + self._client = BabyAIHttpClient(self.env_addr) + try: + self._env_id: int = self._client.create() + except Exception as e: + logging.error(f"[BabyAIEnv] Failed to create env on server: {e}") + self._env_id = -1 + + self._action_list: Optional[List[str]] = None + self._mission: str = "" + self._server_halted: bool = False + self.finished: bool = False + self._init_flag: bool = False + self.episode_return: float = 0.0 + self._last_reward: float = 0.0 + self._timestep: int = 0 + + self.observation_space = gym.spaces.Dict() + self.action_space = gym.spaces.Discrete(self.max_action_num) + self.reward_space = gym.spaces.Box(low=-np.inf, high=np.inf, shape=(1,), dtype=np.float32) + + def prepare_obs(self, obs: str, return_str: bool = False) -> Dict[str, Any]: + raw_obs_text = obs + available_actions = self._action_list if self._action_list else [] + + full_obs = f"{obs}\nValid actions: {available_actions}" + full_obs_str = copy.deepcopy(full_obs) + + if not return_str: + tokenized = BabyAIEnv.tokenizer( + [full_obs], truncation=True, padding="max_length", max_length=self.max_seq_len + ) + obs_attn_mask = tokenized['attention_mask'] + full_obs = np.array(tokenized['input_ids'][0], dtype=np.int32) + + if len(available_actions) == 0: + action_mask = [1] + [0] * (self.max_action_num - 1) + elif len(available_actions) <= self.max_action_num: + action_mask = [1] * len(available_actions) + [0] * (self.max_action_num - len(available_actions)) + else: + action_mask = [1] * self.max_action_num + action_mask = np.array(action_mask, dtype=np.int8) + + if return_str: + result = { + 'observation': full_obs, + 'action_mask': action_mask, + 'valid_actions': available_actions, + 'raw_obs_text': raw_obs_text, + } + if self.for_unizero: + result['to_play'] = -1 + result['timestep'] = self._timestep + return result + else: + result = { + 'observation': full_obs, + 'obs_attn_mask': obs_attn_mask, + 'action_mask': action_mask, + 'valid_actions': available_actions, + 'raw_obs_text': raw_obs_text, + } + if self.for_unizero: + result['to_play'] = -1 + result['timestep'] = self._timestep + return result + + def reset(self, return_str: bool = False) -> Dict[str, Any]: + # aligned with ScalingInter-RL: multi-task cycling + if self.data_idx_list is not None: + if self._is_collect: + self.data_idx = _random.choice(self.data_idx_list) + else: + with BabyAIEnv._eval_cycle_lock: + self.data_idx = self.data_idx_list[ + BabyAIEnv._eval_cycle_counter % len(self.data_idx_list) + ] + BabyAIEnv._eval_cycle_counter += 1 + + if self._server_halted: + try: + self._env_id = self._client.create() + self._server_halted = False + except Exception: + pass + + try: + resp = self._client.reset(self._env_id, self.data_idx) + except Exception as e: + logging.warning(f"[BabyAIEnv] reset failed: {e}") + self._server_halted = True + self._action_list = [] + self._mission = "" + self.finished = False + self._init_flag = True + self.episode_return = 0.0 + self._last_reward = 0.0 + self._timestep = 0 + return self.prepare_obs("[Server unreachable]", return_str) + + obs_text = resp.get('observation', '') + self._mission = _parse_mission(obs_text) + self._action_list = _parse_available_actions(obs_text) + if not self.use_high_level_actions: + self._action_list = list(ATOMIC_ACTIONS) + raw_obs = _strip_actions_suffix(obs_text) + + self.finished = False + self._init_flag = True + self._server_halted = False + self.episode_return = 0.0 + self._last_reward = 0.0 + self._timestep = 0 + + return self.prepare_obs(raw_obs, return_str) + + def step(self, action: Union[int, np.ndarray, str], return_str: bool = False) -> BaseEnvTimestep: + if self._server_halted: + dummy_obs = self.prepare_obs("[Server halted]", return_str) + info = {'action_str': 'noop', 'abnormal': True, 'eval_episode_return': self.episode_return, 'score': self.episode_return} + return BaseEnvTimestep(dummy_obs, 0.0, True, info) + + if isinstance(action, str): + action_str = action + else: + if isinstance(action, np.ndarray): + action = int(action) + try: + action_str = self._action_list[action] + except (IndexError, TypeError): + if self._action_list and len(self._action_list) > 0: + action = int(np.random.choice(len(self._action_list))) + action_str = self._action_list[action] + else: + action_str = "check available actions" + + try: + resp = self._client.step(self._env_id, action_str) + except Exception as e: + logging.warning(f"[BabyAIEnv] step failed on '{action_str}': {e}") + self._server_halted = True + dummy_obs = self.prepare_obs("[Server halted]", return_str) + info = {'action_str': action_str, 'abnormal': True, 'eval_episode_return': self.episode_return, 'score': self.episode_return} + return BaseEnvTimestep(dummy_obs, 0.0, True, info) + + obs_text = resp.get('observation', '') + reward_from_server = float(resp.get('reward', 0.0)) + done = bool(resp.get('done', False)) + + step_reward = reward_from_server - self._last_reward + self._last_reward = reward_from_server + self.episode_return = reward_from_server + + self._timestep += 1 + self._action_list = _parse_available_actions(obs_text) + if not self.use_high_level_actions: + self._action_list = list(ATOMIC_ACTIONS) + raw_obs = _strip_actions_suffix(obs_text) + + if self._timestep >= self.max_steps: + done = True + + processed_obs = self.prepare_obs(raw_obs, return_str) + # aligned with ScalingInter-RL: include task identity for per-level eval logging + info = { + 'action_str': action_str, + 'score': self.episode_return, + 'data_idx': self.data_idx, + 'level_id': self.data_idx % 40 + 1, + } + + if done: + self.finished = True + info['eval_episode_return'] = self.episode_return + + return BaseEnvTimestep(processed_obs, step_reward, done, info) + + def seed(self, seed: int, dynamic_seed: bool = True) -> None: + self._seed = seed + + def close(self) -> None: + self._init_flag = False + if hasattr(self, '_client') and self._client is not None: + self._client.close(self._env_id) + + def __repr__(self) -> str: + return "LightZero BabyAI Env" + + @staticmethod + def create_collector_env_cfg(cfg: Dict[str, Any]) -> List[Dict[str, Any]]: + collector_env_num = cfg.pop('collector_env_num') + cfg = copy.deepcopy(cfg) + cfg['is_collect'] = True + # aligned with ScalingInter-RL: use train task list for collector + if 'train_data_idx_list' in cfg and cfg['train_data_idx_list'] is not None: + cfg['data_idx_list'] = cfg['train_data_idx_list'] + return [cfg for _ in range(collector_env_num)] + + @staticmethod + def create_evaluator_env_cfg(cfg: Dict[str, Any]) -> List[Dict[str, Any]]: + evaluator_env_num = cfg.pop('evaluator_env_num') + cfg = copy.deepcopy(cfg) + cfg['is_collect'] = False + # aligned with ScalingInter-RL: use eval task list for evaluator + if 'eval_data_idx_list' in cfg and cfg['eval_data_idx_list'] is not None: + cfg['data_idx_list'] = cfg['eval_data_idx_list'] + return [cfg for _ in range(evaluator_env_num)] diff --git a/zoo/babyai/priorzero/envs/test_babyai_env.py b/zoo/babyai/priorzero/envs/test_babyai_env.py new file mode 100644 index 000000000..f874dd343 --- /dev/null +++ b/zoo/babyai/priorzero/envs/test_babyai_env.py @@ -0,0 +1,20 @@ +from zoo.babyai.priorzero.envs.babyai_env import BabyAIEnv +cfg = dict(env_addr='http://127.0.0.1:8000', data_idx=0, max_steps=64, + max_action_num=20, tokenizer_path='/mnt/shared-storage-user/puyuan/xiongjyu/models/bge-base-en-v1.5', + max_seq_len=512, for_unizero=True, use_high_level_actions=True, + collector_env_num=1, evaluator_env_num=1) +env = BabyAIEnv(cfg) +obs = env.reset(return_str=True) +print('=== RESET ===') +print('mission:', obs.get('raw_obs_text', '')[:200]) +print('valid_actions:', obs['valid_actions']) +print('action_mask:', obs['action_mask'][:10]) +print('num_actions:', sum(obs['action_mask'])) +for i in range(5): + action = obs['valid_actions'][0] if obs['valid_actions'] else 'check available actions' + ts = env.step(action, return_str=True) + print(f'step {i}: action={action}, reward={ts.reward:.4f}, done={ts.done}') + if ts.done: break + obs = ts.obs +env.close() +print('=== DONE ===') \ No newline at end of file diff --git a/zoo/babyai/priorzero/scripts/run_priorzero_ddp.sh b/zoo/babyai/priorzero/scripts/run_priorzero_ddp.sh new file mode 100644 index 000000000..edad635e6 --- /dev/null +++ b/zoo/babyai/priorzero/scripts/run_priorzero_ddp.sh @@ -0,0 +1,68 @@ +#!/bin/bash +set -x + +cd /mnt/shared-storage-user/puyuan/code/LightZero/zoo/babyai/priorzero +export PYTHONPATH=/mnt/shared-storage-user/puyuan/code/LightZero:$PYTHONPATH + +# ============================================================================ +# PREREQUISITE: Start BabyAI server FIRST +# cd /path/to/AgentGym-RL/AgentGym/agentenv-babyai +# python -m agentenv_babyai.launch --port 8000 +# ============================================================================ + +# 1. Training environment parameters +CUDA_DEVICES="0,1,2,3" +NPROC_PER_NODE=4 + +# CUDA_DEVICES="1,2,3" +# NPROC_PER_NODE=3 + + +# CUDA_DEVICES="2,3" +# NPROC_PER_NODE=2 + +# CUDA_DEVICES="0,1" +# NPROC_PER_NODE=2 + +# MASTER_PORT=24554 +MASTER_PORT=24555 + + +# 2. BabyAI-specific parameters +AGENTGYM_SERVER_ADDR="http://127.0.0.1:8000" +USE_HIGH_LEVEL=true # true = server high-level actions, false = 7 atomic actions + +# 3. Model parameters (aligned with ScalingInter-RL: Qwen2.5-7B, multi-task on 40 levels) +LLM_MODEL="qwen2.5-7b" # "qwen2.5-0.5b" "qwen2.5-1.5b" "qwen2.5-3b" "qwen2.5-7b" +SEED=0 + +USE_COT=true +LOG_DIR="./data_priorzero/babyai/run_logs" +mkdir -p "${LOG_DIR}" + +CURRENT_TIME=$(date +"%Y%m%d_%H%M%S") +LOG_FILE="${LOG_DIR}/log_multitask_${LLM_MODEL}_${CURRENT_TIME}.txt" + +# 4. Environment variables +export CUDA_VISIBLE_DEVICES="${CUDA_DEVICES}" +export PYTHONFAULTHANDLER=1 +export TORCH_DISTRIBUTED_DEBUG=OFF +export NCCL_DEBUG=WARN + +# 5. Build command +CMD_ARGS="--env_id babyai --env_addr ${AGENTGYM_SERVER_ADDR} --model ${LLM_MODEL} --seed ${SEED}" + +if [ "${USE_COT}" = true ]; then + CMD_ARGS="${CMD_ARGS} --use_cot" +fi + +if [ "${USE_HIGH_LEVEL}" = false ]; then + CMD_ARGS="${CMD_ARGS} --use_low_level_actions" +fi + +torchrun \ + --nproc_per_node="${NPROC_PER_NODE}" \ + --master-port="${MASTER_PORT}" \ + ./src/priorzero_entry_sync_ddp.py \ + ${CMD_ARGS} \ + 2>&1 | tee "${LOG_FILE}" diff --git a/zoo/babyai/priorzero/scripts/test_1gpu.sh b/zoo/babyai/priorzero/scripts/test_1gpu.sh new file mode 100644 index 000000000..41df3b829 --- /dev/null +++ b/zoo/babyai/priorzero/scripts/test_1gpu.sh @@ -0,0 +1,7 @@ +#!/bin/bash +set -x + +cd /mnt/shared-storage-user/puyuan/code/LightZero/zoo/babyai/priorzero +export PYTHONPATH=/mnt/shared-storage-user/puyuan/code/LightZero:$PYTHONPATH + +torchrun --nproc_per_node=1 --master-port=24554 ./src/priorzero_entry_sync_ddp.py --quick_test --env_addr http://127.0.0.1:8000 --data_idx 0 --model qwen2.5-3b diff --git a/zoo/babyai/priorzero/src/priorzero_config.py b/zoo/babyai/priorzero/src/priorzero_config.py new file mode 100644 index 000000000..f68b5ccfb --- /dev/null +++ b/zoo/babyai/priorzero/src/priorzero_config.py @@ -0,0 +1,450 @@ +import os +from typing import Dict, Tuple, Optional, Any +from easydict import EasyDict +import torch.distributed as dist +from dataclasses import dataclass, field + +# ============================================================================ +# Model Configuration Presets (shared with Jericho version) +# ============================================================================ +MODEL_CONFIGS = { + "qwen2.5-0.5b": { + "model_name_or_path": "/mnt/shared-storage-user/puyuan/xiongjyu/models/Qwen2.5-0.5B-Instruct", + "vllm_tensor_parallel_size": 1, + "gpu_memory_utilization": 0.2, + "description": "Qwen2.5-0.5B-Instruct (smallest, fastest)", + }, + "qwen2.5-1.5b": { + "model_name_or_path": "/mnt/shared-storage-user/puyuan/xiongjyu/models/Qwen2.5-1.5B-Instruct", + "vllm_tensor_parallel_size": 1, + "gpu_memory_utilization": 0.2, + "description": "Qwen2.5-1.5B-Instruct (balanced performance)", + }, + "qwen2.5-3b": { + "model_name_or_path": "/mnt/shared-storage-user/puyuan/xiongjyu/models/Qwen2.5-3B-Instruct", + "vllm_tensor_parallel_size": 1, + "gpu_memory_utilization": 0.25, + "description": "Qwen2.5-3B-Instruct (better quality)", + }, + "qwen2.5-7b": { + "model_name_or_path": "/mnt/shared-storage-user/puyuan/xiongjyu/models/Qwen2.5-7B-Instruct", + "vllm_tensor_parallel_size": 2, + # "vllm_tensor_parallel_size": 1, + + "gpu_memory_utilization": 0.35, + "description": "Qwen2.5-7B-Instruct (high quality, needs 2+ GPUs)", + }, + "qwen2.5-14b": { + "model_name_or_path": "/mnt/shared-storage-user/puyuan/model/Qwen2.5-14B-Instruct", + "vllm_tensor_parallel_size": 4, + "gpu_memory_utilization": 0.5, + "description": "Qwen2.5-14B-Instruct (best quality, needs 4+ GPUs)", + }, +} + +def get_available_models(): + return list(MODEL_CONFIGS.keys()) + +def get_model_config(model_key: str) -> Dict: + if model_key not in MODEL_CONFIGS: + available = ", ".join(get_available_models()) + raise ValueError(f"Unknown model key: {model_key}\nAvailable models: {available}") + return MODEL_CONFIGS[model_key] + + +@dataclass +class PriorZeroLLMConfig: + model_name_or_path: str = "Qwen2.5-3B-Instruct" + local_rank: int = -1 + enable_rft: bool = True + enable_world_model: bool = True + train_mode_dict: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "mode": "full", + "lora_r": 16, + "lora_alpha": 32, + "lora_dropout": 0.05, + "lora_bias": "none", + "lora_target_modules": ( + "q_proj", "k_proj", "v_proj", "o_proj", + "gate_proj", "up_proj", "down_proj", + ), + })) + + train_schedule: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "alternate": True, + "wm_update_iters": 500, + "llm_update_iters": 100, + "start_phase": "wm", + "wm_warmup_updates": 0, + })) + + llm_prior_temperature: float = 1.0 + mcts_root_logits_dict: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "mode": "llm_plus_wm_logits", + "plus_method": "fixed", + "wm_weight": 0.5, + "llm_max_weight": 0.7, + "llm_min_weight": 0.3, + "max_envsteps": 1e5, + })) + eval_dict: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "world_model": True, + "world_model_llm_prior": True, + "llm_prior": True, + "wm_eval_freq": 500, + "llm_eval_freq": 50, + # env-step-based eval frequency (preferred over iter-based when > 0) + "wm_eval_freq_envsteps": 0, # 0 = disabled, falls back to wm_eval_freq + "llm_eval_freq_envsteps": 0, # 0 = disabled, falls back to llm_eval_freq + "save_llm_cot": True, # 是否在 eval trajectory JSON 中保存 CoT/prompt/LLM prior 等详细信息 + })) + + attn_implementation: str = "flash_attention_2" + history_length: int = 10 + use_cot: bool = True + cot_weight: float = 0.1 + + user_prompt_dict: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "history_with_reward": True, + "observation_with_valid_actions": True, + })) + + # Total context budget consumed by line 662 of priorzero_datafactory.py as + # `max_length = prompt_max_len - generate_max_len - 20`; BabyAI obs typically ≤ 512 tokens. + prompt_max_len: int = 4096 + generate_max_len: int = 512 + bf16: bool = True + + enable_vllm: bool = True + enable_prefix_caching: bool = False + use_cuda_ipc: bool = False + enable_vllm_is_correction: bool = False + vllm_is_truncated_threshold: Tuple[float, float] = (0.5, 5.0) + use_mispo: bool = False + mispo_token_truncated_threshold: Tuple[float, float] = (0.5, 2.0) + mispo_traj_truncated_threshold: Tuple[float, float] = (0.8, 1.2) + + vllm_sync_backend: str = "nccl" + vllm_tensor_parallel_size: int = 1 + gpu_memory_utilization: float = 0.3 + vllm_enable_sleep: bool = True + temperature: float = 1.0 + top_p: float = 0.95 + seed: int = 0 + + reduction: str = "mean" + + deepspeed_enable_sleep: bool = True + zero_stage: int = 2 + gradient_checkpointing: bool = False + gradient_checkpointing_use_reentrant: bool = False + max_norm: float = 1.0 + ds_tensor_parallel_size: int = 1 + + train_batch_size: int = 128 + micro_train_batch_size: int = 2 + max_rollout_staleness: int = 1 + + learning_rate: float = 1e-6 + adam_betas: Tuple[float, float] = (0.9, 0.95) + weight_decay: float = 0.01 + lr_scheduler: str = "cosine_with_min_lr" + lr_warmup_ratio: float = 0.03 + max_steps: int = int(1e4) + policy_loss_type: str = "ppo" + reward_func: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "format_reward": True, + "format_param": EasyDict({"format_weight": 0.3}), + })) + advantage_type: str = "advantage_global_batch_norm" + eps_clip_low_high: Tuple[float, float] = (0.2, 0.2) + rft_kl_coef: float = 0.1 + entropy_loss_coef: float = 0.001 + kl_estimator: str = "k3" + # KL early stopping: skip remaining micro-batches when ref_kl exceeds this threshold + # kl_early_stop_threshold: float = 0.1 + kl_early_stop_threshold: float = 0.0 + + + llm_save_freq: int = 1000 + save_path: str = "" + + value_norm_cfg: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + 'enable_stability_optimizer': True, + 'value_norm_init_momentum': 0.9, + 'value_norm_final_momentum': 0.99, + 'value_norm_warmup_steps': 100, + 'value_norm_clip_percentile': 0.95, + 'value_norm_clip_method': "soft", + "value_norm_history_size": 1000, + })) + + +def get_priorzero_config( + env_id: str = 'babyai', + seed: int = 0, + exp_name: str = None, + use_cot: bool = True, + model_key: Optional[str] = "qwen2.5-3b", + multi_gpu: bool = False, + env_addr: str = 'http://127.0.0.1:8000', + use_high_level_actions: bool = True, +) -> Tuple[EasyDict, EasyDict]: + + action_space_size = 20 # upper bound for dynamic action space + max_steps = 20 # aligned with ScalingInter-RL babyai_train.sh (max_rounds=20) + wm_encoder_option = 'legacy' + wm_model_name = '/mnt/shared-storage-user/puyuan/xiongjyu/models/bge-base-en-v1.5' + + # Aligned with ScalingInter-RL (HF: AgentGym/AgentGym-RL-Data-ID, train/babyai_train.json). + # ScalingInter-RL trains on 18 out of 40 BabyAI levels (810 items, 45 seeds per level). + # BabyAI level mapping: level_id = data_idx % 40 + 1, seed = data_idx // 40. + # Using seed=0 (data_idx = level_id - 1) for PriorZero since it re-samples each episode. + _SCALING_INTER_RL_LEVELS = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 19, 20, 21, 30, 31, 33, 36] + train_data_idx_list = [lvl - 1 for lvl in _SCALING_INTER_RL_LEVELS] # 18 levels, seed=0 + eval_data_idx_list = [lvl - 1 for lvl in _SCALING_INTER_RL_LEVELS] # same 18 levels for eval + + collector_env_num = 1 + # Set evaluator_env_num == n_evaluator_episode so each env runs exactly one episode + # (covers every eval level once and avoids the buggy `n_episode > env_num` refill path). + evaluator_env_num = len(eval_data_idx_list) + evaluator_env_num = 4 + + + n_episode = collector_env_num + n_evaluator_episode = len(eval_data_idx_list) # 18 episodes to cover all eval levels + + + # only for debug + # evaluator_env_num = 2 + # n_evaluator_episode = 2 + + num_unroll_steps = 10 + infer_context_length = 4 + game_segment_length = 50 + num_layers = 2 + embed_dim = 768 + replay_ratio = 0.1 + batch_size = 64 + # collect_num_simulations = 50 + collect_num_simulations = 25 + eval_num_simulations = 50 + + # only for debug + # collect_num_simulations = 2 + # eval_num_simulations = 2 + + # replay_buffer_size = int(3e5) + replay_buffer_size = int(5e5) + + + env_config = dict( + stop_value=int(1e6), + max_steps=max_steps, + observation_shape=512, + env_id=env_id, + env_addr=env_addr, + train_data_idx_list=train_data_idx_list, # aligned with ScalingInter-RL + eval_data_idx_list=eval_data_idx_list, # aligned with ScalingInter-RL + use_high_level_actions=use_high_level_actions, + for_unizero=True, + tokenizer_path=wm_model_name, + max_action_num=action_space_size, + max_seq_len=512, + collector_env_num=collector_env_num, + evaluator_env_num=evaluator_env_num, + n_evaluator_episode=n_evaluator_episode, # aligned with ScalingInter-RL: cover all eval tasks + manager=dict(shared_memory=False), + ) + policy_config = dict( + type='priorzero', + multi_gpu=multi_gpu, + use_wandb=False, + learn=dict( + learner=dict( + hook=dict(save_ckpt_after_iter=1000000), + ), + ), + model=dict( + observation_shape=512, + action_space_size=action_space_size, + encoder_option=wm_encoder_option, + encoder_url=wm_model_name, + model_type="mlp", + continuous_action_space=False, + norm_type="LN", + world_model_cfg=dict( + norm_type="LN", + final_norm_option_in_head="LayerNorm", + final_norm_option_in_encoder="LayerNorm", + predict_latent_loss_type='mse', + policy_entropy_weight=5e-2, + continuous_action_space=False, + max_blocks=num_unroll_steps, + max_tokens=2 * num_unroll_steps, + context_length=2 * infer_context_length, + device="cuda", + action_space_size=action_space_size, + num_layers=num_layers, + num_heads=24, + embed_dim=embed_dim, + obs_type="text", + env_num=max(collector_env_num, evaluator_env_num), + decode_loss_mode=None, + latent_recon_loss_weight=0, + task_embed_option=None, + moe_in_transformer=False, + multiplication_moe_in_transformer=False, + game_segment_length=game_segment_length, + ) + ), + update_per_collect=None, + num_segments=collector_env_num, + action_type="varied_action_space", + model_path=None, + num_unroll_steps=num_unroll_steps, + reanalyze_ratio=0, + replay_ratio=replay_ratio, + batch_size=batch_size, + learning_rate=3e-4, + weight_decay=1e-4, + cos_lr_scheduler=False, + fixed_temperature_value=0.25, + manual_temperature_decay=False, + n_episode=n_episode, + train_start_after_envsteps=0, + replay_buffer_size=replay_buffer_size, + # eval_freq=int(3e4), + eval_freq=int(2e4), + collector_env_num=collector_env_num, + evaluator_env_num=evaluator_env_num, + buffer_reanalyze_freq=1 / 1000000, + reanalyze_batch_size=160, + reanalyze_partition=0.75, + device='cuda', + collect_num_simulations=collect_num_simulations, + eval_num_simulations=eval_num_simulations, + game_segment_length=game_segment_length, + off_policy_degree=0, + enable_async_eval=False, + optim_type='AdamW', + grad_clip_value=10.0, + value_loss_weight=0.25, + policy_loss_weight=1.0, + reward_loss_weight=1.0, + use_adaptive_entropy_weight=False, + adaptive_entropy_alpha_lr=1e-4, + use_encoder_clip_annealing=False, + encoder_clip_anneal_type='cosine', + encoder_clip_start_value=30.0, + encoder_clip_end_value=10.0, + encoder_clip_anneal_steps=100000, + use_priority=False, + priority_prob_alpha=0.6, + priority_prob_beta=0.4, + ) + + llm_config = PriorZeroLLMConfig(use_cot=use_cot) + + model_config = get_model_config(model_key) + llm_config.model_name_or_path = model_config["model_name_or_path"] + llm_config.vllm_tensor_parallel_size = model_config["vllm_tensor_parallel_size"] + llm_config.gpu_memory_utilization = model_config["gpu_memory_utilization"] + + if exp_name is None: + # aligned with ScalingInter-RL: multi-task across 18 levels + if llm_config.enable_rft: + exp_name = ( + f"data_priorzero/babyai/llm_rft/priorzero_multitask_18levels_{model_key}_train_{llm_config.train_mode_dict.mode}/" + f"useCot_{llm_config.use_cot}_alternate_{llm_config.train_schedule.alternate}/" + f"mcts_{llm_config.mcts_root_logits_dict.mode}_staleness_{llm_config.max_rollout_staleness}_tbs_{llm_config.train_batch_size}_use-mispo-{llm_config.use_mispo}_seed{seed}" + ) + else: + exp_name = ( + f"data_priorzero/babyai/llm_frozen/priorzero_multitask_18levels_{model_key}_" + f"train_{llm_config.train_mode_dict.mode}" + f"useCot_{llm_config.use_cot}_seed{seed}" + ) + + priorzero_config = dict( + env=env_config, + policy=policy_config, + exp_name=exp_name, + seed=seed + ) + create_config = dict( + env=dict( + type="babyai", + import_names=["zoo.babyai.priorzero.envs.babyai_env"], + ), + env_manager=dict(type="base"), + policy=dict( + type="priorzero", + import_names=["zoo.jericho.priorzero.src.priorzero_policy"], + ), + collector=dict( + type="priorzero_segment", + import_names=["zoo.jericho.priorzero.src.priorzero_collector"], + ), + evaluator=dict( + type="priorzero", + import_names=["zoo.jericho.priorzero.src.priorzero_evaluator"], + ), + replay_buffer=dict( + type='game_buffer_muzero', + import_names=['lzero.mcts.buffer.game_buffer_muzero'], + ), + ) + main_config = EasyDict(priorzero_config) + create_config = EasyDict(create_config) + + train_level_ids = [idx % 40 + 1 for idx in train_data_idx_list] + eval_level_ids = [idx % 40 + 1 for idx in eval_data_idx_list] + import logging + logging.getLogger("priorzero.main").info( + f"[Config] model={model_key} | {len(train_data_idx_list)} train levels | {len(eval_data_idx_list)} eval levels | high_level={use_high_level_actions}" + ) + + return main_config, create_config, llm_config + + +def get_priorzero_debug_config( + env_id: str = 'babyai', + # seed: int = 0, + seed: int = 1, + exp_name: str = None, + use_cot: bool = True, + model_key: Optional[str] = "qwen2.5-3b", + env_addr: str = 'http://127.0.0.1:8000', + use_high_level_actions: bool = True, +) -> EasyDict: + + main_config, create_config, llm_config = get_priorzero_config( + env_id=env_id, seed=seed, exp_name=exp_name, use_cot=use_cot, + model_key=model_key, env_addr=env_addr, + use_high_level_actions=use_high_level_actions, + ) + max_steps = 20 + batch_size = 8 + collect_num_simulations = 2 + eval_num_simulations = 2 + num_layers = 1 + game_segment_length = 50 + + llm_config.train_batch_size = 8 + llm_config.micro_train_batch_size = 4 + llm_config.train_schedule.wm_update_iters = 2 + llm_config.train_schedule.llm_update_iters = 1 + llm_config.eval_dict.wm_eval_freq = 2 + llm_config.eval_dict.llm_eval_freq = 1 + + main_config.env.max_steps = max_steps + main_config.policy.model.world_model_cfg.num_layers = num_layers + main_config.policy.model.world_model_cfg.game_segment_length = game_segment_length + main_config.policy.batch_size = batch_size + main_config.policy.collect_num_simulations = collect_num_simulations + main_config.policy.eval_num_simulations = eval_num_simulations + main_config.policy.update_per_collect = 2 + main_config.policy.game_segment_length = game_segment_length + + return main_config, create_config, llm_config diff --git a/zoo/babyai/priorzero/src/priorzero_datafactory.py b/zoo/babyai/priorzero/src/priorzero_datafactory.py new file mode 100644 index 000000000..243aaa894 --- /dev/null +++ b/zoo/babyai/priorzero/src/priorzero_datafactory.py @@ -0,0 +1,87 @@ +import importlib.util +from pathlib import Path +from typing import List, Tuple, Optional + +_jericho_df_path = str( + Path(__file__).resolve().parent.parent.parent.parent + / "jericho" / "priorzero" / "src" / "priorzero_datafactory.py" +) +_spec = importlib.util.spec_from_file_location("jericho_datafactory", _jericho_df_path) +_jericho_mod = importlib.util.module_from_spec(_spec) +_spec.loader.exec_module(_jericho_mod) +JerichoDataProcessor = _jericho_mod.DataProcessor + + +class DataProcessor(JerichoDataProcessor): + """BabyAI-specific DataProcessor with grid-world appropriate prompts.""" + + def get_system_prompt(self): + parts = [ + "You are an expert agent navigating a BabyAI grid-world environment. " + "You are placed in rooms and must accomplish goals by choosing optimal actions.", + "", + "Available action types:", + "- turn left / turn right / move forward: basic movement", + "- go to : navigate to a specific object", + "- pick up : pick up an object", + "- go through : go through an open door", + "- toggle and go through : open and go through a closed/locked door (locked doors require a matching color key)", + "- toggle: open/close a door directly in front of you", + "", + "Your goal is to complete the given task efficiently to maximize your score.", + "", + "OUTPUT FORMAT:", + ] + if self.use_cot: + parts.append( + "You MUST produce exactly TWO parts in the following order:\n" + "1. Reasoning: Analyze the current observation, your position, nearby objects, " + "and which action best progresses toward the goal.\n" + "2. Action: The final chosen action (must be one of the valid actions).\n" + "Strict Format Example:\n" + "Reasoning: \n" + "Action: " + ) + else: + parts.append( + "Output exactly one line starting with 'Action:'.\n" + "Example:\n" + "Action: " + ) + return "\n".join(parts) + + def get_user_prompt(self, history=None, current_obs=None, valid_actions=None): + prompt_parts = [] + user_prompt_dict = self.args.user_prompt_dict + + if history and len(history) > 0: + prompt_parts.append("=== ACTION HISTORY ===") + for i, (obs, action, reward) in enumerate(history, start=1): + prompt_parts.append(f"Step {i}:") + prompt_parts.append(f"Observation: {obs.strip()}") + prompt_parts.append(f"Action: {action.strip()}") + if user_prompt_dict.history_with_reward: + prompt_parts.append(f"Reward: {reward}") + prompt_parts.append("") + + prompt_parts.append("=== CURRENT OBSERVATION ===") + prompt_parts.append(current_obs.strip()) + + if user_prompt_dict.observation_with_valid_actions: + if valid_actions and len(valid_actions) > 0: + actions_str = ", ".join([f"'{act}'" for act in valid_actions]) + prompt_parts.append(f"\n[Valid Actions]\nChoose from: {actions_str}") + + prompt_parts.append("\n=== INSTRUCTION ===") + if self.use_cot: + prompt_parts.append( + "Analyze the observation and provide your response:\n" + "Reasoning: \n" + "Action: " + ) + else: + prompt_parts.append( + "Choose the best action:\n" + "Action: " + ) + return "\n".join(prompt_parts) diff --git a/zoo/babyai/priorzero/src/priorzero_entry_sync_ddp.py b/zoo/babyai/priorzero/src/priorzero_entry_sync_ddp.py new file mode 100644 index 000000000..e862be6bd --- /dev/null +++ b/zoo/babyai/priorzero/src/priorzero_entry_sync_ddp.py @@ -0,0 +1,349 @@ +import sys +import os +import logging +from pathlib import Path + +# Add Jericho PriorZero src to path for shared modules +_jericho_src = str(Path(__file__).resolve().parent.parent.parent.parent / "jericho" / "priorzero" / "src") +# Local src dir first so priorzero_config resolves to BabyAI version +_local_src = str(Path(__file__).resolve().parent) +sys.path.insert(0, _jericho_src) +sys.path.insert(0, _local_src) + +import asyncio +from functools import partial +from typing import Tuple, Optional, List + +import torch +import torch.distributed as dist +import wandb + +from ding.config import compile_config, save_config +from ding.envs import create_env_manager, get_vec_env_setting +from ding.policy import create_policy +from ding.utils import set_pkg_seed, get_rank, get_world_size +from ding.worker import create_buffer, BaseLearner +from tensorboardX import SummaryWriter + +from priorzero_config import ( + get_priorzero_config, + get_priorzero_debug_config, + get_available_models, +) +from priorzero_collector import PriorZeroCollector +from priorzero_evaluator import PriorZeroEvaluator +from priorzero_policy import * +from lzero.mcts.buffer.game_buffer_priorzero import PriorZeroGameBufferOptimized +from utils import dump_dataclass_cfg_py, setup_priorzero_logging + +from lzero.entry.utils import calculate_update_per_collect + +_log_main = logging.getLogger("priorzero.main") +_log_train = logging.getLogger("priorzero.train") +_log_eval = logging.getLogger("priorzero.eval") + +def prepare_unizero(rank, cfg, create_cfg, llm_cfg, seed): + cfg = compile_config(cfg, seed=seed, auto=True, create_cfg=create_cfg) + env_fn, collector_env_cfg, evaluator_env_cfg = get_vec_env_setting(cfg.env) + collector_env = create_env_manager(cfg.env.manager, [partial(env_fn, cfg=c) for c in collector_env_cfg]) + evaluator_env = create_env_manager(cfg.env.manager, [partial(env_fn, cfg=c) for c in evaluator_env_cfg]) + + collector_env.seed(seed) + evaluator_env.seed(seed, dynamic_seed=False) + + policy = create_policy(cfg.policy, enable_field=['learn', 'collect', 'eval'], exp_name=cfg.exp_name, llm_cfg=llm_cfg) + if cfg.policy.model_path is not None: + _log_main.info(f"Loading pretrained model from {cfg.policy.model_path}") + policy.learn_mode.load_state_dict(torch.load(cfg.policy.model_path, map_location=cfg.policy.device)) + + os.makedirs(f'./{cfg.exp_name}/log/', exist_ok=True) + tb_logger = SummaryWriter(os.path.join(f'./{cfg.exp_name}/log/', 'serial')) if get_rank() == 0 else None + + learner = BaseLearner( + cfg.policy.learn.learner, + policy.learn_mode, + tb_logger, + exp_name=cfg.exp_name + ) + + replay_buffer = PriorZeroGameBufferOptimized(cfg.policy) + + collector = PriorZeroCollector( + env=collector_env, + policy=policy.collect_mode, + llm_config=llm_cfg, + tb_logger=tb_logger, + exp_name=cfg.exp_name, + policy_config=cfg.policy, + ) + + evaluator = PriorZeroEvaluator( + n_evaluator_episode=cfg.env.n_evaluator_episode, + stop_value=cfg.env.stop_value, + env=evaluator_env, + policy=policy.eval_mode, + tb_logger=tb_logger, + exp_name=cfg.exp_name, + policy_config=cfg.policy, + llm_config=llm_cfg, + ) + learner.call_hook('before_run') + _log_main.info("Policy, Learner, Collector, Evaluator created") + + return cfg, replay_buffer, tb_logger, policy, collector, evaluator, learner + +def all_gather_cmd(world_size, obj) -> List: + if world_size <= 1: + return [obj] + lst = [None] * dist.get_world_size() + dist.all_gather_object(lst, obj) + return lst + +def train_priorzero( + cfg: dict, + create_cfg: dict, + llm_cfg, + seed: int = 0, + max_train_iter: int = int(1e6), + max_env_step: Optional[int] = int(1e10), + enable_profile: bool = False +): + rank = int(os.environ.get("RANK", "0")) + from strategy.deepspeed import get_strategy, torch_dist_barrier_and_cuda_sync + strategy = get_strategy(llm_cfg) + + strategy.setup_distributed() + world_size = getattr(strategy, "world_size", 1) + + cfg, replay_buffer, tb_logger, policy, collector, evaluator, learner = prepare_unizero( + rank=rank, cfg=cfg, create_cfg=create_cfg, llm_cfg=llm_cfg, seed=seed + ) + batch_size = cfg.policy.batch_size + + # Initialize structured logging after exp_name is known + setup_priorzero_logging(cfg.exp_name, rank) + _log_main.info(f"=== PriorZero Training Start | rank={rank}/{world_size} | exp={cfg.exp_name} ===") + + if rank == 0: + dump_dataclass_cfg_py(llm_cfg, path=f"{cfg.exp_name}/llm_cfg.py") + # Save config snapshot + import yaml + config_path = os.path.join(cfg.exp_name, "run_logs", "config.yaml") + with open(config_path, "w") as f: + yaml.dump({"llm_cfg": str(llm_cfg), "policy_cfg": str(cfg.policy)}, f, default_flow_style=False) + llm_cfg.save_path = f'./{cfg.exp_name}/llm_ckpt/' + + from utils import Profiler + prof = Profiler(log_interval=10, stats_file=f'./{cfg.exp_name}/log/profiler.txt', enable_profile=enable_profile) + + _log_main.info("Initializing LLM Actor...") + set_pkg_seed(seed + rank, use_cuda=True) + + from models.actor import PolicyModel, ReferenceModel + if llm_cfg.rft_kl_coef > 0: + ref_model = ReferenceModel(strategy=strategy, pretrain=llm_cfg.model_name_or_path) + else: + ref_model = None + + from vllm_utils.vllm_engine import create_vllm_engine + vllm_engine = create_vllm_engine( + tensor_parallel_size=llm_cfg.vllm_tensor_parallel_size, + pretrain=llm_cfg.model_name_or_path, + enable_prefix_caching=llm_cfg.enable_prefix_caching, + max_model_len=llm_cfg.prompt_max_len + llm_cfg.generate_max_len, + gpu_memory_utilization=llm_cfg.gpu_memory_utilization, + vllm_enable_sleep=llm_cfg.vllm_enable_sleep, + ) + _log_main.info("vLLM engine created") + + from priorzero_datafactory import DataProcessor + data_processor = DataProcessor( + rank=rank, world_size=world_size, vllm_engine=vllm_engine, + strategy=strategy, model_path=llm_cfg.model_name_or_path, + exp_name=cfg.exp_name if rank == 0 else None, + ) + collector.data_processor = data_processor + collector.prof = prof + evaluator.data_processor = data_processor + + policy_model = PolicyModel( + strategy=strategy, pretrain=llm_cfg.model_name_or_path, + vllm_engine=vllm_engine, max_steps=llm_cfg.max_steps + ) + from priorzero_trainer import PriorZeroLLMTrainer + trainer = PriorZeroLLMTrainer( + cfg=llm_cfg, pretrain=llm_cfg.model_name_or_path, + strategy=strategy, vllm_engine=vllm_engine, + policy_model=policy_model, reference_model=ref_model, + exp_name=cfg.exp_name if rank == 0 else None, + tb_logger=tb_logger if rank == 0 else None, + llm_save_freq=llm_cfg.llm_save_freq + ) + + torch_dist_barrier_and_cuda_sync() + train_schedule = llm_cfg.train_schedule + train_alternate = train_schedule["alternate"] + current_phase = None + if train_alternate: + current_phase = train_schedule["start_phase"] + last_wm_train_iter = 0 + last_llm_train_iter = 0 + + _log_eval.info("=== Initial Evaluation ===") + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.wake_up() + evaluator.eval(wm_train_iter=0, llm_train_iter=0, phase=current_phase, env_step=collector.envstep) + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.sleep() + torch_dist_barrier_and_cuda_sync() + + _log_main.info(f"=== Training Loop Start | phase={current_phase} ===") + while True: + if collector.envstep >= max_env_step or learner.train_iter >= max_train_iter: + break + + if learner.train_iter != 0 and evaluator.should_eval(wm_train_iter=learner.train_iter, llm_train_iter=policy_model.train_iter, phase=current_phase, env_step=collector.envstep): + _log_eval.info(f"=== Eval | wm_iter={learner.train_iter} llm_iter={policy_model.train_iter} phase={current_phase} envstep={collector.envstep} ===") + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.wake_up() + evaluator.eval(wm_train_iter=learner.train_iter, llm_train_iter=policy_model.train_iter, phase=current_phase, env_step=collector.envstep) + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.sleep() + + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.wake_up() + + new_data = collector.collect(train_iter=learner.train_iter, policy_kwargs={'temperature': 0.25, 'epsilon': 0.0}, phase=current_phase) + data_processor.get_llm_output_log(wm_train_iter=learner.train_iter, llm_train_iter=policy_model.train_iter) + + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.sleep() + + replay_buffer.push_game_segments(new_data) + replay_buffer.remove_oldest_data_to_fit() + num_of_transitions = replay_buffer.get_num_of_transitions() + + torch_dist_barrier_and_cuda_sync() + + if llm_cfg.enable_world_model and (not train_alternate or (train_alternate and current_phase == "wm")): + if not (num_of_transitions > batch_size): + _log_train.warning(f"[WM] Data insufficient: buffer={num_of_transitions} < batch={batch_size}") + cmd = 0 + else: + cmd = 1 + if min(all_gather_cmd(world_size=world_size, obj=cmd)) == 0: + continue + + update_per_collect = calculate_update_per_collect(cfg, new_data, world_size=world_size) + _log_train.info(f"[WM] Iter {learner.train_iter} | updates={update_per_collect} | buffer={num_of_transitions}") + + for i in range(update_per_collect): + with prof.block("train_world_model", rank=rank): + train_data = replay_buffer.sample(batch_size, policy) + train_data.append(learner.train_iter) + log_vars = learner.train(train_data, collector.envstep) + if cfg.policy.use_priority: + replay_buffer.update_priority(train_data, log_vars[0]['value_priority_orig']) + policy.recompute_pos_emb_diff_and_clear_cache() + if llm_cfg.enable_rft and train_alternate and learner.train_iter - last_wm_train_iter >= train_schedule["wm_update_iters"]: + current_phase = "llm" + last_wm_train_iter = learner.train_iter + replay_buffer.mark_latest_transitions_consumed() + _log_main.info(f"=== Phase Switch: WM -> LLM | wm_iter={learner.train_iter} ===") + continue + + if llm_cfg.enable_rft and (not train_alternate or (train_alternate and current_phase == "llm")): + new_num_of_transitions = replay_buffer.get_num_of_transitions() - replay_buffer.last_pos_in_transition + _log_train.info(f"[LLM] Total={num_of_transitions} | New={new_num_of_transitions}") + + with prof.block("fetch_latest_batch", rank=rank): + priorzero_batch = replay_buffer.fetch_latest_batch(batch_size=-1, policy=policy) + torch.cuda.empty_cache() + + with prof.block("train_llm", rank=rank): + llm_need_sample_cnt = llm_cfg.train_batch_size * llm_cfg.max_rollout_staleness // world_size + flag, train_samples = data_processor.make_llm_train_samples(priorzero_batch, ddp=True, max_samples=llm_need_sample_cnt) + + if not flag: + local_llm_ready = 0 + else: + local_llm_ready = 1 + gathered_llm_ready = all_gather_cmd(world_size=world_size, obj=local_llm_ready) + + if min(gathered_llm_ready) == 0: + _log_train.debug(f"Skip LLM training: not all ranks ready. flags={gathered_llm_ready}") + continue + + trainer.train_batch(train_samples, collect_env_steps=collector.envstep) + replay_buffer.mark_latest_transitions_consumed() + + torch_dist_barrier_and_cuda_sync() + if llm_cfg.enable_world_model and train_alternate and trainer.global_step - last_llm_train_iter >= train_schedule["llm_update_iters"]: + current_phase = "wm" + last_llm_train_iter = trainer.global_step + data_processor.clear_statis() + _log_main.info(f"=== Phase Switch: LLM -> WM | llm_iter={trainer.global_step} ===") + +def main(): + import argparse + import requests as req + + parser = argparse.ArgumentParser(description='PriorZero BabyAI Training') + parser.add_argument('--env_id', type=str, default='babyai', help='Environment ID') + parser.add_argument('--env_addr', type=str, default='http://127.0.0.1:8000', help='BabyAI server address') + parser.add_argument('--use_high_level_actions', action='store_true', default=True, help='Use server high-level actions') + parser.add_argument('--use_low_level_actions', action='store_true', default=False, help='Use 7 atomic actions') + parser.add_argument('--seed', type=int, default=0, help='Random seed') + parser.add_argument('--max_iter', type=int, default=int(1e6), help='Max training iterations') + parser.add_argument('--quick_test', action='store_true', default=False, help='Use debug config') + parser.add_argument('--model', type=str, default="qwen2.5-3b", choices=get_available_models()) + parser.add_argument('--enable_profile', action='store_true', default=False) + parser.add_argument('--use_cot', action='store_true', default=False) + args = parser.parse_args() + + use_high_level = not args.use_low_level_actions + model_key = args.model + + # args.seed = 2 + # args.seed = 3 + + + + rank = int(os.environ.get("RANK", "0")) + + if rank == 0: + try: + r = req.get(f"{args.env_addr}/", timeout=5) + assert r.status_code == 200, f"Server returned status {r.status_code}" + except Exception as e: + raise RuntimeError( + f"BabyAI server not reachable at {args.env_addr}: {e}\n" + f"Start it first: cd /AgentGym/agentenv-babyai && python -m agentenv_babyai.launch --port 8000" + ) + print(f"[PriorZero] model={model_key} | server={args.env_addr} | 18 levels | seed={args.seed} | cot={args.use_cot}") + + if args.quick_test: + main_cfg, create_cfg, llm_cfg = get_priorzero_debug_config( + args.env_id, args.seed, use_cot=args.use_cot, + exp_name='data_priorzero/babyai/priorzero_debug_multitask', + model_key=model_key, env_addr=args.env_addr, + use_high_level_actions=use_high_level, + ) + else: + main_cfg, create_cfg, llm_cfg = get_priorzero_config( + args.env_id, args.seed, use_cot=args.use_cot, + model_key=model_key, multi_gpu=True, + env_addr=args.env_addr, + use_high_level_actions=use_high_level, + ) + + train_priorzero( + main_cfg, create_cfg, llm_cfg, + seed=args.seed, max_train_iter=args.max_iter, + enable_profile=args.enable_profile, + ) + + +if __name__ == "__main__": + os.environ['TOKENIZERS_PARALLELISM'] = 'false' + main() diff --git a/zoo/box2d/lunarlander/config/lunarlander_image_unizero_config.py b/zoo/box2d/lunarlander/config/lunarlander_image_unizero_config.py new file mode 100644 index 000000000..4751e58f7 --- /dev/null +++ b/zoo/box2d/lunarlander/config/lunarlander_image_unizero_config.py @@ -0,0 +1,254 @@ +import sys +import os +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '..', '..', '..', '..'))) + +from easydict import EasyDict + +# ============================================================== +# Debug mode: set True to print detailed loss/grad diagnostics every 50 train iters +# ============================================================== +debug_mode = True +# ============================================================== +# begin of the most frequently changed config specified by the user +# ============================================================== +collector_env_num = 8 +num_segments = 8 +evaluator_env_num = 3 +num_simulations = 50 +reanalyze_ratio = 0. +update_per_collect = None +replay_ratio = 0.25 +# replay_ratio = 0.1 +max_env_step = int(1e6) +batch_size = 256 +num_unroll_steps = 10 +infer_context_length = 4 +num_layers = 2 +norm_type = 'BN' +game_segment_length = 200 + +buffer_reanalyze_freq = 1/5000000000 +reanalyze_batch_size = 160 +reanalyze_partition = 0.75 + +# debug +# collector_env_num = 2 +# num_segments = 2 +# evaluator_env_num = 2 +# num_simulations = 5 +# batch_size = 2 +# ============================================================== +# end of the most frequently changed config specified by the user +# ============================================================== + +lunarlander_image_unizero_config = dict( + exp_name=f'data_unizero_0422_debug/lunarlander_image_unizero_ns{num_simulations}_upc{update_per_collect}-rr{replay_ratio}_rer{reanalyze_ratio}_H{num_unroll_steps}-infer{infer_context_length}_bs{batch_size}_{norm_type}_seed0', + env=dict( + env_id='LunarLander-v2', + # observation_shape=(3, 64, 64), + observation_shape=(3, 96, 96), + image_size=96, + gray_scale=False, + continuous=False, + manually_discretization=False, + collector_env_num=collector_env_num, + evaluator_env_num=evaluator_env_num, + n_evaluator_episode=evaluator_env_num, + manager=dict(shared_memory=False, ), + collect_max_episode_steps=int(10000), + eval_max_episode_steps=int(10000), + ), + policy=dict( + learn=dict(learner=dict(hook=dict(save_ckpt_after_iter=1000000, ), ), ), + model=dict( + # observation_shape=(3, 64, 64), + observation_shape=(3, 96, 96), + action_space_size=4, + norm_type=norm_type, + # ====== [FIX] support range must cover LunarLander reward/value range (-200 ~ +300) ====== + reward_support_range=(-300., 301., 1.), + value_support_range=(-300., 301., 1.), + num_res_blocks=1, + num_channels=64, + world_model_cfg=dict( + observation_shape=(3, 96, 96), + continuous_action_space=False, + max_blocks=num_unroll_steps, + max_tokens=2 * num_unroll_steps, + context_length=2 * infer_context_length, + device='cuda', + action_space_size=4, + num_layers=num_layers, + num_heads=4, + embed_dim=256, + obs_type='image', + encoder_type='resnet', + group_size=8, + norm_type=norm_type, + env_num=max(collector_env_num, evaluator_env_num), + support_size=601, + # Normalization options + # final_norm_option_in_encoder='LayerNorm', + # final_norm_option_in_obs_head='LayerNorm', + # predict_latent_loss_type='mse', + final_norm_option_in_encoder='SimNorm', + final_norm_option_in_obs_head='SimNorm', + predict_latent_loss_type='group_kl', + # Task embedding (single-task, disabled) + task_embed_option=None, + # MoE (disabled for single-task baseline) + moe_in_transformer=False, + multiplication_moe_in_transformer=False, + # Misc + policy_entropy_weight=5e-3, + num_simulations=num_simulations, + game_segment_length=game_segment_length, + rotary_emb=False, + latent_recon_loss_weight=0., + perceptual_loss_weight=0., + decode_loss_mode=None, + use_priority=False, + use_normal_head=True, + use_softmoe_head=False, + use_moe_head=False, + # optim_type='AdamW_mix_lr_wdecay', + optim_type='AdamW', + ), + ), + model_path=None, + num_unroll_steps=num_unroll_steps, + cuda=True, + game_segment_length=game_segment_length, + update_per_collect=update_per_collect, + batch_size=batch_size, + # optim_type='AdamW_mix_lr_wdecay', + # weight_decay=1e-2, + optim_type='AdamW', + # weight_decay=1e-2, + learning_rate=0.0001, + piecewise_decay_lr_scheduler=False, + num_simulations=num_simulations, + reanalyze_ratio=reanalyze_ratio, + num_segments=num_segments, + replay_ratio=replay_ratio, + replay_buffer_size=int(1e6), + collector_env_num=collector_env_num, + evaluator_env_num=evaluator_env_num, + # ====== [FIX] grad clip: 20 -> 5, prevent gradient explosion ====== + grad_clip_value=5, + # ====== [FIX] Priority Experience Replay ====== + # use_priority=True, + use_priority=False, + priority_prob_alpha=1, + priority_prob_beta=1, + # ====== [FIX] Adaptive entropy weight ====== + # use_adaptive_entropy_weight=True, + use_adaptive_entropy_weight=False, + adaptive_entropy_alpha_lr=1e-4, + target_entropy_start_ratio=0.98, + target_entropy_end_ratio=0.7, + target_entropy_decay_steps=100000, + # ====== [FIX] Encoder-clip annealing ====== + # use_encoder_clip_annealing=True, + use_encoder_clip_annealing=False, + encoder_clip_anneal_type='cosine', + encoder_clip_start_value=30.0, + encoder_clip_end_value=10.0, + encoder_clip_anneal_steps=100000, + # ====== [FIX] Label smoothing ====== + policy_ls_eps_start=0.0, + policy_ls_eps_end=0.01, + policy_ls_eps_decay_steps=50000, + label_smoothing_eps=0.1, + # ====== Monitor ====== + monitor_norm_freq=10000, + eval_freq=int(5e3), + td_steps=5, + train_start_after_envsteps=0, + use_augmentation=True, + # manual_temperature_decay=False, + manual_temperature_decay=True, + threshold_training_steps_for_final_temperature=int(5e4), + # ============= Reanalyze ============= + buffer_reanalyze_freq=buffer_reanalyze_freq, + reanalyze_batch_size=reanalyze_batch_size, + reanalyze_partition=reanalyze_partition, + ), +) +lunarlander_image_unizero_config = EasyDict(lunarlander_image_unizero_config) +main_config = lunarlander_image_unizero_config + +lunarlander_image_unizero_create_config = dict( + env=dict( + type='lunarlander_image', + import_names=['zoo.box2d.lunarlander.envs.lunarlander_image_env'], + ), + env_manager=dict(type='subprocess'), + policy=dict( + type='unizero', + import_names=['lzero.policy.unizero'], + ), +) +lunarlander_image_unizero_create_config = EasyDict(lunarlander_image_unizero_create_config) +create_config = lunarlander_image_unizero_create_config + +if __name__ == "__main__": + import logging + logging.basicConfig(level=logging.DEBUG if debug_mode else logging.INFO, + format='[%(asctime)s][%(name)s][%(levelname)s] %(message)s') + # NOTE: 日志文件请通过 shell 重定向实现,例如: + # python lunarlander_image_unizero_config.py 2>&1 | tee /mnt/shared-storage-user/puyuan/code/LightZero/data_unizero_0422_debug/logs/train_$(date +%Y%m%d_%H%M%S).log + + # ====== Debug mode: monkey-patch to print diagnostics ====== + if debug_mode: + from ding.worker import BaseLearner + _original_train = BaseLearner.train + + def _debug_train(self, data, envstep=-1): + log_vars = _original_train(self, data, envstep) + if log_vars and self.train_iter % 50 == 0: + d = log_vars[0] if isinstance(log_vars, list) else log_vars + def _fmt(v): + if v == 'N/A': + return 'N/A' + try: + return f'{float(v):.4f}' + except (TypeError, ValueError): + return str(v) + logging.info( + f"[DEBUG] iter={self.train_iter} envstep={envstep} | " + f"total_loss={_fmt(d.get('weighted_total_loss', 'N/A'))} | " + f"policy={_fmt(d.get('policy_loss', 'N/A'))} | " + f"value={_fmt(d.get('value_loss', 'N/A'))} | " + f"reward={_fmt(d.get('reward_loss', 'N/A'))} | " + f"obs={_fmt(d.get('obs_loss', 'N/A'))} | " + f"entropy={_fmt(d.get('policy_entropy', 'N/A'))} | " + f"target_policy_entropy={_fmt(d.get('target_policy_entropy', 'N/A'))} | " + f"grad_norm={_fmt(d.get('total_grad_norm_before_clip_wm', 'N/A'))} | " + f"lr={_fmt(d.get('cur_lr_world_model', 'N/A'))} | " + f"target_reward={_fmt(d.get('target_reward', 'N/A'))} | " + f"target_value={_fmt(d.get('target_value', 'N/A'))} | " + f"dormant_enc={d.get('analysis/dormant_ratio_encoder', 'N/A')} | " + f"dormant_tf={d.get('analysis/dormant_ratio_transformer', 'N/A')} | " + f"latent_l2={d.get('analysis/latent_state_l2_norms', 'N/A')} | " + f"GPU={_fmt(d.get('Current_GPU', 'N/A'))}GB" + ) + return log_vars + + BaseLearner.train = _debug_train + + from lzero.worker import MuZeroSegmentCollector + _original_output_log = MuZeroSegmentCollector._output_log + + def _debug_output_log(self, train_iter): + _original_output_log(self, train_iter) + logging.info( + f"[DEBUG][Collector] total_envstep={self._total_envstep_count} " + f"total_episode={self._total_episode_count}" + ) + + MuZeroSegmentCollector._output_log = _debug_output_log + + # ====== Train ====== + from lzero.entry import train_unizero_segment + train_unizero_segment([main_config, create_config], seed=0, model_path=main_config.policy.model_path, max_env_step=max_env_step) diff --git a/zoo/box2d/lunarlander/envs/lunarlander_image_env.py b/zoo/box2d/lunarlander/envs/lunarlander_image_env.py new file mode 100644 index 000000000..9b54a9fe1 --- /dev/null +++ b/zoo/box2d/lunarlander/envs/lunarlander_image_env.py @@ -0,0 +1,157 @@ +""" +Image-based LunarLander Environment for PriorZero VLM + +Wraps the standard LunarLander-v2 to produce image observations (3, 64, 64) +instead of vector observations, enabling VLM-based prior generation. +""" +import copy +from typing import List, Dict + +import cv2 +import gymnasium as gym +import numpy as np +from ding.torch_utils import to_ndarray +from ding.utils import ENV_REGISTRY +from easydict import EasyDict + +from zoo.box2d.lunarlander.envs.lunarlander_env import LunarLanderEnv + + +@ENV_REGISTRY.register('lunarlander_image') +class LunarLanderImageEnv(LunarLanderEnv): + """ + Image-based LunarLander environment. + + Replaces the 8-dim vector observation with a (3, 64, 64) RGB image + rendered from the environment. Everything else (actions, rewards, done) + remains identical to the base LunarLanderEnv. + """ + + config = dict( + env_id="LunarLander-v2", + save_replay_gif=False, + replay_path_gif=None, + replay_path=None, + act_scale=False, + collect_max_episode_steps=int(1000), + eval_max_episode_steps=int(1000), + image_size=64, + ) + + @classmethod + def default_config(cls) -> EasyDict: + cfg = EasyDict(copy.deepcopy(cls.config)) + cfg.cfg_type = cls.__name__ + 'Dict' + return cfg + + def __init__(self, cfg: dict) -> None: + super().__init__(cfg) + self._image_size = cfg.get('image_size', 64) + + def _render_image_obs(self) -> np.ndarray: + """Render the environment and return a (3, H, W) float32 image scaled to [0, 1].""" + frame = self._env.render() # (H, W, 3) RGB uint8 + # Resize to target size + frame = cv2.resize(frame, (self._image_size, self._image_size), interpolation=cv2.INTER_AREA) + # HWC -> CHW, scale to [0, 1] float32 (consistent with Atari env scale=True) + frame = np.transpose(frame, (2, 0, 1)).astype(np.float32) / 255.0 + return frame + + def reset(self) -> Dict[str, np.ndarray]: + if not self._init_flag: + self._env = gym.make(self._cfg.env_id, render_mode="rgb_array") + self._observation_space = gym.spaces.Box( + low=0, high=1, shape=(3, self._image_size, self._image_size), dtype=np.float32 + ) + self._action_space = self._env.action_space + self._reward_space = gym.spaces.Box( + low=self._env.reward_range[0], high=self._env.reward_range[1], shape=(1,), dtype=np.float32 + ) + self._init_flag = True + + if hasattr(self, '_seed') and hasattr(self, '_dynamic_seed') and self._dynamic_seed: + np_seed = 100 * np.random.randint(1, 1000) + self._seed = self._seed + np_seed + self._env.reset(seed=self._seed) + elif hasattr(self, '_seed'): + self._env.reset(seed=self._seed) + else: + self._env.reset() + + self._eval_episode_return = 0.0 + self._timestep = 0 + if self._save_replay_gif: + self._frames = [] + + # Render image observation + obs_image = self._render_image_obs() + action_mask = np.ones(4, 'int8') + obs = { + 'observation': obs_image, + 'action_mask': action_mask, + 'to_play': -1, + 'timestep': self._timestep, + } + return obs + + def step(self, action: np.ndarray): + from ding.envs import BaseEnvTimestep + + if isinstance(action, np.ndarray) and action.shape == (1,): + action = action.item() + elif not isinstance(action, np.ndarray): + action = int(action) + if self._save_replay_gif: + self._frames.append(self._env.render()) + + _, rew, terminated, truncated, info = self._env.step(action) + done = terminated or truncated + self._timestep += 1 + + # Render image observation + obs_image = self._render_image_obs() + action_mask = np.ones(4, 'int8') + obs = { + 'observation': obs_image, + 'action_mask': action_mask, + 'to_play': -1, + 'timestep': self._timestep, + } + + self._eval_episode_return += rew + if done: + info['eval_episode_return'] = self._eval_episode_return + info['score'] = self._eval_episode_return + if self._save_replay_gif: + import os + from datetime import datetime + if not os.path.exists(self._replay_path_gif): + os.makedirs(self._replay_path_gif) + timestamp = datetime.now().strftime("%Y%m%d%H%M%S") + path = os.path.join( + self._replay_path_gif, + f'{self._env_id}_episode_{self._save_replay_count}_seed{self._seed}_{timestamp}.gif' + ) + self.display_frames_as_gif(self._frames, path) + self._save_replay_count += 1 + + obs = to_ndarray(obs) + rew = to_ndarray(rew).astype(np.float32) + return BaseEnvTimestep(obs, rew, done, info) + + @staticmethod + def create_collector_env_cfg(cfg: dict) -> List[dict]: + collector_env_num = cfg.pop('collector_env_num') + cfg = copy.deepcopy(cfg) + cfg.max_episode_steps = cfg.collect_max_episode_steps + return [cfg for _ in range(collector_env_num)] + + @staticmethod + def create_evaluator_env_cfg(cfg: dict) -> List[dict]: + evaluator_env_num = cfg.pop('evaluator_env_num') + cfg = copy.deepcopy(cfg) + cfg.max_episode_steps = cfg.eval_max_episode_steps + return [cfg for _ in range(evaluator_env_num)] + + def __repr__(self) -> str: + return "LightZero LunarLander Image Env." diff --git a/zoo/jericho/bkp/fused_unizero_config.py b/zoo/jericho/bkp/fused_unizero_config.py deleted file mode 100644 index 095e25c09..000000000 --- a/zoo/jericho/bkp/fused_unizero_config.py +++ /dev/null @@ -1,127 +0,0 @@ -# fused_unizero_config.py - -import os -from easydict import EasyDict - -def get_priorzero_config(env_id: str = 'zork1.z5', seed: int = 0) -> EasyDict: - """ - Generates the configuration for the PriorZero algorithm, merging UniZero and LLM settings. - """ - # ============================================================== - # 1. UniZero Base Configurations - # ============================================================== - action_space_size, max_steps = 20, 100 # Default for Jericho, can be overridden - - # World Model Encoder (can be different from the main policy LLM) - wm_encoder_option = 'legacy' - if wm_encoder_option == 'legacy': - wm_model_name = 'BAAI/bge-base-en-v1.5' - else: - wm_model_name = 'Qwen/Qwen2-0.5B' # A smaller model for the world model encoder - - jericho_unizero_config = dict( - env=dict( - stop_value=int(1e6), - max_steps=max_steps, - observation_shape=768, # Embedding dimension - max_action_num=action_space_size, - tokenizer_path=wm_model_name, - game_path=f"./zoo/jericho/envs/z-machine-games-master/jericho-game-suite/{env_id}", - collector_env_num=8, - evaluator_env_num=5, - n_evaluator_episode=5, - manager=dict(shared_memory=False), - ), - policy=dict( - # This section now primarily configures the World Model and MCTS - model=dict( - observation_shape=768, - action_space_size=action_space_size, - encoder_option=wm_encoder_option, - encoder_url=wm_model_name, - model_type="mlp", - world_model_cfg=dict( - final_norm_option_in_obs_head='LayerNorm', - final_norm_option_in_encoder='LayerNorm', - predict_latent_loss_type='mse', - policy_entropy_weight=5e-3, - continuous_action_space=False, - max_blocks=10, # num_unroll_steps - max_tokens=20, - context_length=8, # 2 * infer_context_length - device="cuda", - action_space_size=action_space_size, - num_layers=4, - num_heads=12, - embed_dim=768, - obs_type="text", - env_num=8, - decode_loss_mode='None', - latent_recon_loss_weight=0.1, - ), - ), - # MCTS settings - num_simulations=50, - root_dirichlet_alpha=0.3, - root_noise_weight=0.25, - # World Model training settings - batch_size=64, - num_unroll_steps=10, - td_steps=5, - learning_rate=3e-4, # LR for World Model - weight_decay=1e-4, - # Replay Buffer settings - replay_buffer_size=int(5e4), - replay_ratio=0.25, - # Other RL settings - eval_freq=int(1e3), - train_start_after_envsteps=2000, - ), - ) - - # ============================================================== - # 2. LLM Policy (ORZ-style) Configurations - # ============================================================== - llm_policy_config = dict( - # Model path for the main LLM policy - pretrain="Qwen/Qwen2.5-7B", - # vLLM settings for efficient inference - vllm_num_engines=jericho_unizero_config['env']['collector_env_num'], - vllm_tensor_parallel_size=1, - gpu_memory_utilization=0.7, - # LLM Policy training settings (RFT/PPO) - llm_learning_rate=1e-6, - llm_weight_decay=0.01, - # Prompting - prompt_max_len=4096, - generate_max_len=512, - ) - - # Add LLM config to the main policy config - jericho_unizero_config['policy']['llm_policy_cfg'] = llm_policy_config - - # ============================================================== - # 3. Create Config for DI-engine - # ============================================================== - create_config = dict( - env=dict( - type="jericho", - import_names=["zoo.jericho.envs.jericho_env"], - ), - env_manager=dict(type="base"), - # We will create a custom policy class `PriorZeroPolicy` - policy=dict( - type="priorzero", # Register a new policy type - import_names=["your_project.policy.priorzero_policy"], # Path to your custom policy - ), - ) - - # ============================================================== - # 4. Final Touches - # ============================================================== - main_config = EasyDict(jericho_unizero_config) - create_config = EasyDict(create_config) - - main_config.exp_name = f"data_lz/priorzero/{env_id}_qwen7b_seed{seed}" - - return main_config, create_config \ No newline at end of file diff --git a/zoo/jericho/bkp/fused_unizero_entry.py b/zoo/jericho/bkp/fused_unizero_entry.py deleted file mode 100644 index 31f52f1a1..000000000 --- a/zoo/jericho/bkp/fused_unizero_entry.py +++ /dev/null @@ -1,293 +0,0 @@ -# fused_unizero_entry.py - -import asyncio -import os -from functools import partial -from typing import Tuple, Optional, List, Dict - -import ray -import torch -import numpy as np -from ding.config import compile_config -from ding.envs import create_env_manager, get_vec_env_setting -from ding.policy import create_policy -from ding.utils import set_pkg_seed, get_rank -from tensorboardX import SummaryWriter -from loguru import logger - -# Import necessary components from LightZero/UniZero -from lzero.entry.utils import log_buffer_memory_usage, calculate_update_per_collect -from lzero.policy import visit_count_temperature -from lzero.worker import MuZeroSegmentCollector as UniZeroCollector -from lzero.worker import MuZeroEvaluator as Evaluator -from lzero.mcts import UniZeroGameBuffer # The replay buffer -from ding.worker import BaseLearner - -# Import ORZ/vLLM components for LLM inference -from vllm import AsyncLLMEngine, SamplingParams -from vllm.engine.arg_utils import AsyncEngineArgs - -# --- Custom Components for PriorZero --- - -class PriorZeroCollector(UniZeroCollector): - """ - Custom Collector for PriorZero. - It uses an LLM for policy priors at the MCTS root and a World Model for search. - """ - def __init__(self, env, policy, tb_logger, exp_name, policy_config, vllm_engine: AsyncLLMEngine): - super().__init__(env, policy, tb_logger, exp_name, policy_config) - self.vllm_engine = vllm_engine - self.llm_policy_cfg = policy_config.llm_policy_cfg - logger.info("PriorZeroCollector initialized with vLLM engine.") - - async def _async_get_llm_prior(self, states: List[str]) -> List[Dict]: - """ Asynchronously gets policy priors from the LLM. """ - prompts = [] - for state in states: - instruction = ( - "You are an expert player in a text-based adventure game. " - "Based on the history, think step-by-step and propose a ranked list of the best actions to take next. " - "Your goal is to maximize the score.\n\n" - f"=== History ===\n{state}\n\n" - "=== Analysis and Ranked Actions (e.g., 1. take key 2. look) ===" - ) - # NOTE: Assuming the policy model uses the same tokenizer as ORZ - prompts.append(self._policy.llm_policy_model_tokenizer.apply_chat_template( - [{"role": "user", "content": instruction}], tokenize=False - )) - - sampling_params = SamplingParams( - temperature=1.0, top_p=1.0, max_tokens=self.llm_policy_cfg.generate_max_len, stop=["==="] - ) - - request_ids = [f"collect_{self._collect_count}_{i}" for i in range(len(prompts))] - results_generator = self.vllm_engine.generate(prompts, sampling_params, request_ids) - - llm_outputs = [] - async for result in results_generator: - llm_outputs.append(result) - - # Sort results back to original order - llm_outputs.sort(key=lambda r: int(r.request_id.split('_')[-1])) - return llm_outputs - - @override - async def collect(self, n_segment: Optional[int] = None, train_iter: int = 0, policy_kwargs: Optional[dict] = None) -> List[Dict]: - """ - Asynchronous data collection method. - """ - # This is a simplified version of the collection loop. - # A full implementation would handle multiple segments and episodes. - - # Get current states and valid actions from all parallel envs - # This part requires modification in the env_manager to be async or batched - current_obs = self._env.ready_obs - states = [obs['raw_obs'] for obs in current_obs.values()] - valid_actions_list = [obs['action_mask'] for obs in current_obs.values()] # Assuming this format - - # 1. Get policy priors from LLM asynchronously - llm_outputs = await self._async_get_llm_prior(states) - - # The rest of the logic is inside _forward_collect of the policy - # We need to pass the LLM priors to it. - policy_kwargs = policy_kwargs or {} - policy_kwargs['llm_outputs'] = llm_outputs - - # The original `collect` is synchronous. We are calling the internal `_collect` logic here. - # This part needs significant re-engineering to fit the async model. - # For this blueprint, we assume the policy's forward pass can handle this. - - # The original call is synchronous, we are showing the conceptual flow - # In a real implementation, `self._policy._forward_collect` would need to be async - # and handle the interaction loop. - - # For now, let's just say the policy's collect function is now async - # and we await it. This implies deep changes in the policy class itself. - - # Conceptual: The policy's `_forward_collect` will now: - # a. Parse llm_outputs to create root priors. - # b. Run MCTS using the world model. - # c. Sample actions and step the environments. - # d. Return the collected game segments. - - # This is a placeholder for the complex interaction logic. - # The key is that the `collect` method is now `async`. - logger.info("Conceptual async collect step completed.") - # In a real system, this would return collected data segments. - # We will mock this by returning an empty list, assuming data is pushed to buffer inside policy. - - # Let's simulate one step and data push for demonstration - # This logic would actually be inside the policy/collector loop - mock_game_segments = [] - for i in range(len(states)): - # Mock MCTS result - mcts_policy = np.ones(len(valid_actions_list[i])) / len(valid_actions_list[i]) - action = np.random.choice(len(valid_actions_list[i])) - # Mock env step - # self._env.step(...) - mock_game_segments.append({'state': states[i], 'action': action, 'mcts_policy': mcts_policy}) - - return mock_game_segments - - -class PriorZeroLearner(BaseLearner): - """ - Custom Learner for PriorZero. - Trains both the World Model and the LLM Policy. - """ - def _init_learn(self): - # This method is called by BaseLearner's __init__ - self.world_model = self._policy.world_model - self.llm_policy_model = self._policy.llm_policy_model - - # Optimizer for World Model - self.world_model_optimizer = torch.optim.AdamW( - self.world_model.parameters(), - lr=self._cfg.learning_rate, # From UniZero config - weight_decay=self._cfg.weight_decay - ) - - # Optimizer for LLM Policy Model - # This assumes the LLM is loaded and managed by the policy - self.llm_policy_optimizer = torch.optim.AdamW( - self.llm_policy_model.parameters(), - lr=self._cfg.llm_policy_cfg.llm_learning_rate, - weight_decay=self._cfg.llm_policy_cfg.llm_weight_decay - ) - - def _forward(self, data: List[Dict]) -> Dict[str, any]: - """ - The main training step. - """ - # --- 1. World Model Update --- - # Prepare batch for world model (as in UniZero) - # This is a complex data transformation step - # wm_batch = self._policy.prepare_data_for_wm(data) - world_model_loss_info = self.world_model.compute_loss({}) # Mocked call - wm_loss = world_model_loss_info.loss_total - - self.world_model_optimizer.zero_grad() - wm_loss.backward() - self.world_model_optimizer.step() - - # --- 2. LLM Policy Update (RFT) --- - # Prepare batch for LLM policy (instruction tuning format) - # llm_batch = self._policy.prepare_data_for_llm(data) - - # For simplicity, we'll implement a Behavior Cloning (SFT) loss - # The LLM should predict the MCTS policy - # In a real PPO setup, this would be much more complex - - # Conceptual SFT loss: - # llm_inputs = self.tokenizer(llm_batch['prompts'], return_tensors='pt', padding=True) - # target_logits = self.tokenizer(llm_batch['targets'], return_tensors='pt', padding=True).input_ids - # outputs = self.llm_policy_model(**llm_inputs, labels=target_logits) - # llm_loss = outputs.loss - llm_loss = torch.tensor(0.1, requires_grad=True) # Mock loss - - self.llm_policy_optimizer.zero_grad() - llm_loss.backward() - self.llm_policy_optimizer.step() - - return { - 'wm_loss': wm_loss.item(), - 'llm_loss': llm_loss.item(), - } - -async def train_priorzero( - input_cfg: Tuple[dict, dict], - seed: int = 0, - max_env_step: Optional[int] = int(1e10), -) -> None: - """ - Asynchronous training entry for PriorZero. - """ - cfg, create_cfg = input_cfg - cfg = compile_config(cfg, seed=seed, auto=True, create_cfg=create_cfg) - - # Initialize Ray - if not ray.is_initialized(): - ray.init() - - # 1. Create vLLM Engine as a Ray Actor (like in ORZ) - engine_args = AsyncEngineArgs( - model=cfg.policy.llm_policy_cfg.pretrain, - tensor_parallel_size=cfg.policy.llm_policy_cfg.vllm_tensor_parallel_size, - gpu_memory_utilization=cfg.policy.llm_policy_cfg.gpu_memory_utilization, - worker_use_ray=True, - ) - vllm_engine = AsyncLLMEngine.from_engine_args(engine_args) - logger.info("vLLM Engine created successfully.") - - # 2. Create Environment and Policy - env_fn, collector_env_cfg, evaluator_env_cfg = get_vec_env_setting(cfg.env) - collector_env = create_env_manager(cfg.env.manager, [partial(env_fn, cfg=c) for c in collector_env_cfg]) - evaluator_env = create_env_manager(cfg.env.manager, [partial(env_fn, cfg=c) for c in evaluator_env_cfg]) - collector_env.seed(seed) - evaluator_env.seed(seed, dynamic_seed=False) - set_pkg_seed(seed, use_cuda=True) - - # This will create our custom PriorZeroPolicy - policy = create_policy(cfg.policy, enable_field=['learn', 'collect', 'eval']) - - # 3. Create Custom Worker Components - tb_logger = SummaryWriter(os.path.join(f'./{cfg.exp_name}/log/', 'serial')) - - # Pass the vLLM engine to the collector - collector = PriorZeroCollector( - env=collector_env, - policy=policy.collect_mode, - tb_logger=tb_logger, - exp_name=cfg.exp_name, - policy_config=cfg.policy, - vllm_engine=vllm_engine - ) - - # The learner needs to be our custom one - learner = PriorZeroLearner(cfg.policy.learn.learner, policy.learn_mode, tb_logger, exp_name=cfg.exp_name) - evaluator = Evaluator(eval_freq=cfg.policy.eval_freq, n_evaluator_episode=cfg.env.n_evaluator_episode, - stop_value=cfg.env.stop_value, env=evaluator_env, policy=policy.eval_mode, - tb_logger=tb_logger, exp_name=cfg.exp_name, policy_config=cfg.policy) - - replay_buffer = UniZeroGameBuffer(cfg.policy) - - # --- Main Asynchronous Training Loop --- - learner.call_hook('before_run') - - while collector.envstep < max_env_step: - log_buffer_memory_usage(learner.train_iter, replay_buffer, tb_logger) - - # Collect experience asynchronously - collect_kwargs = {'temperature': visit_count_temperature(trained_steps=learner.train_iter, **cfg.policy)} - new_data = await collector.collect(train_iter=learner.train_iter, policy_kwargs=collect_kwargs) - - replay_buffer.push_game_segments(new_data) - - # Train models if buffer is ready - if collector.envstep > cfg.policy.train_start_after_envsteps: - update_per_collect = calculate_update_per_collect(cfg, new_data) - for i in range(update_per_collect): - train_data = replay_buffer.sample(cfg.policy.batch_size, policy) - if not train_data: - break - log_vars = learner.train(train_data, collector.envstep) - - # Log to tensorboard - for k, v in log_vars.items(): - tb_logger.add_scalar(f'train/{k}', v, learner.train_iter) - - # Evaluation - if evaluator.should_eval(learner.train_iter): - stop, reward = evaluator.eval(learner.save_checkpoint, learner.train_iter, collector.envstep) - if stop: - break - - learner.call_hook('after_run') - - -if __name__ == "__main__": - # Get configuration - main_cfg, create_cfg = get_priorzero_config(env_id='zork1.z5') - - # Start the asynchronous training process - asyncio.run(train_priorzero([main_cfg, create_cfg], seed=0)) \ No newline at end of file diff --git a/zoo/jericho/configs/jericho_unizero_config.py b/zoo/jericho/configs/jericho_unizero_config.py index 45da0b81b..41f9c9b62 100644 --- a/zoo/jericho/configs/jericho_unizero_config.py +++ b/zoo/jericho/configs/jericho_unizero_config.py @@ -18,7 +18,7 @@ def main(env_id: str = 'detective.z5', seed: int = 0, max_env_step: int = int(1e """ env_id = 'detective.z5' - collector_env_num: int = 4 # Number of collector environments + collector_env_num: int = 8 # Number of collector environments n_episode = int(collector_env_num) batch_size=64 @@ -39,7 +39,7 @@ def main(env_id: str = 'detective.z5', seed: int = 0, max_env_step: int = int(1e # ------------------------------------------------------------------ # User frequently modified configurations # ------------------------------------------------------------------ - evaluator_env_num: int = 3 # Number of evaluator environments + evaluator_env_num: int = 8 # Number of evaluator environments num_simulations: int = 50 # Number of simulations # Project training parameters @@ -48,7 +48,7 @@ def main(env_id: str = 'detective.z5', seed: int = 0, max_env_step: int = int(1e num_layers: int = 2 # Number of layers in the model replay_ratio: float = 0.1 # Replay ratio for experience replay - embed_dim: int = 512 # Embedding dimension + embed_dim: int = 768 # Embedding dimension # Reanalysis (reanalyze) parameters: # buffer_reanalyze_freq: Frequency of reanalysis (e.g., 1 means reanalyze once per epoch) @@ -150,7 +150,8 @@ def main(env_id: str = 'detective.z5', seed: int = 0, max_env_step: int = int(1e lora_dropout= 0.0, decode_loss_mode=None, # Controls where to compute reconstruction loss: after_backbone, before_backbone, or None. - latent_recon_loss_weight=0.1 + latent_recon_loss_weight=0.1, + game_segment_length=50 ), ), update_per_collect=int(collector_env_num*max_steps*replay_ratio ), # Important for DDP @@ -168,7 +169,7 @@ def main(env_id: str = 'detective.z5', seed: int = 0, max_env_step: int = int(1e n_episode=n_episode, train_start_after_envsteps=0, # TODO: Adjust training start trigger if needed. replay_buffer_size=int(5e5), - eval_freq=int(3e4), + eval_freq=int(300), collector_env_num=collector_env_num, evaluator_env_num=evaluator_env_num, buffer_reanalyze_freq=buffer_reanalyze_freq, @@ -203,7 +204,7 @@ def main(env_id: str = 'detective.z5', seed: int = 0, max_env_step: int = int(1e # Construct experiment name containing key parameters main_config.exp_name = ( - f"data_lz/data_unizero_jericho/bge-base-en-v1.5/{env_id}/uz_gpu_cen{collector_env_num}_rr{replay_ratio}_ftemp025_{env_id[:8]}_ms{max_steps}_ass-{action_space_size}_" + f"data_lz_fixed2/data_unizero_jericho/bge-base-en-v1.5/{env_id}/uz_gpu_cen{collector_env_num}_rr{replay_ratio}_ftemp025_{env_id[:8]}_ms{max_steps}_ass-{action_space_size}_" f"nlayer{num_layers}_embed{embed_dim}_Htrain{num_unroll_steps}-" f"Hinfer{infer_context_length}_bs{batch_size}_seed{seed}" ) diff --git a/zoo/jericho/configs/jericho_unizero_qwen_ddp_config.py b/zoo/jericho/configs/jericho_unizero_qwen_ddp_config.py new file mode 100644 index 000000000..d79879977 --- /dev/null +++ b/zoo/jericho/configs/jericho_unizero_qwen_ddp_config.py @@ -0,0 +1,214 @@ +import os +import argparse +from typing import Any, Dict + +from easydict import EasyDict + + +def main(env_id: str = 'detective.z5', seed: int = 0, max_env_step: int = int(1e6)) -> None: + """ + DDP entry for Jericho UniZero with Qwen2.5-0.5B as the latent world-model backbone. + + Most settings follow jericho_unizero_ddp_config.py in the current priorzero branch. + The only intended algorithmic change is enabling the Qwen backbone + inside UniZero's world model. + """ + gpu_num = int(os.environ.get("WORLD_SIZE", "4")) + collector_env_num: int = int(os.environ.get("COLLECTOR_ENV_NUM", "4")) + n_episode = int(collector_env_num * gpu_num) + + # Keep the observation encoder from the current DDP config. Qwen is used as + # the world-model backbone, not as the text observation encoder. + encoder_option = 'legacy' + model_name: str = '/mnt/afs/niuyazhe/workspace/xiongjyu/models/bge-base-en-v1.5' + batch_size = int(os.environ.get("BATCH_SIZE", str(64 * gpu_num))) + accumulation_steps = 1 + + qwen_backbone_path: str = '/mnt/afs/niuyazhe/workspace/xiongjyu/models/Qwen2.5-0.5B' + + env_configurations = { + 'detective.z5': (12, 100), + 'omniquest.z5': (25, 100), + 'acorncourt.z5': (45, 50), + 'zork1.z5': (55, 500), + } + action_space_size, max_steps = env_configurations.get(env_id, (10, 50)) + max_steps = int(os.environ.get("MAX_STEPS", max_steps)) + + evaluator_env_num: int = int(os.environ.get("EVALUATOR_ENV_NUM", "3")) + num_simulations: int = int(os.environ.get("NUM_SIMULATIONS", "50")) + num_unroll_steps: int = 10 + infer_context_length: int = 4 + + # Qwen2.5-0.5B config: hidden_size=896, layers=24, attention_heads=14, kv_heads=2. + # UniZero's KV cache stores key/value heads, so num_heads is the number of KV heads. + num_layers: int = 24 + replay_ratio: float = 0.1 + embed_dim: int = 896 + num_heads: int = 2 + hidden_size: int = 64 + + buffer_reanalyze_freq: float = 1 / 100000 + reanalyze_batch_size: int = 160 + reanalyze_partition: float = 0.75 + + jericho_unizero_config: Dict[str, Any] = dict( + env=dict( + stop_value=int(1e6), + observation_shape=512, + max_steps=max_steps, + max_action_num=action_space_size, + tokenizer_path=model_name, + max_seq_len=512, + game_path=f"./zoo/jericho/envs/z-machine-games-master/jericho-game-suite/{env_id}", + for_unizero=True, + collector_env_num=collector_env_num, + evaluator_env_num=evaluator_env_num, + n_evaluator_episode=evaluator_env_num, + manager=dict(shared_memory=False), + ), + policy=dict( + multi_gpu=True, + use_wandb=False, + learn=dict( + learner=dict( + hook=dict( + save_ckpt_after_iter=1000000, + ), + ), + ), + accumulation_steps=accumulation_steps, + model=dict( + observation_shape=512, + action_space_size=action_space_size, + encoder_url=model_name, + encoder_option=encoder_option, + model_type="mlp", + continuous_action_space=False, + world_model_cfg=dict( + final_norm_option_in_obs_head='LayerNorm', + final_norm_option_in_encoder='LayerNorm', + predict_latent_loss_type='mse', + policy_entropy_weight=5e-2, + continuous_action_space=False, + max_blocks=num_unroll_steps, + max_tokens=2 * num_unroll_steps, + context_length=2 * infer_context_length, + device="cuda", + action_space_size=action_space_size, + use_qwen_backbone=True, + pretrained_path=qwen_backbone_path, + num_layers=num_layers, + num_heads=num_heads, + embed_dim=embed_dim, + hidden_size=hidden_size, + obs_type="text", + env_num=max(collector_env_num, evaluator_env_num), + task_embed_option=None, + use_task_embed=False, + use_normal_head=True, + use_softmoe_head=False, + use_moe_head=False, + num_experts_in_moe_head=4, + moe_in_transformer=False, + multiplication_moe_in_transformer=False, + n_shared_experts=1, + num_experts_per_tok=1, + num_experts_of_moe_in_transformer=8, + lora_r=0, + lora_alpha=1, + lora_dropout=0.0, + decode_loss_mode=None, + latent_recon_loss_weight=0.1, + game_segment_length=50, + ), + ), + update_per_collect=int(collector_env_num * max_steps * replay_ratio * accumulation_steps), + action_type="varied_action_space", + model_path=None, + num_unroll_steps=num_unroll_steps, + reanalyze_ratio=0, + replay_ratio=replay_ratio, + batch_size=batch_size, + learning_rate=0.0001, + cos_lr_scheduler=False, + fixed_temperature_value=0.25, + manual_temperature_decay=False, + num_simulations=num_simulations, + n_episode=n_episode, + train_start_after_envsteps=0, + replay_buffer_size=int(5e5), + eval_freq=int(3e4), + collector_env_num=collector_env_num, + evaluator_env_num=evaluator_env_num, + buffer_reanalyze_freq=buffer_reanalyze_freq, + reanalyze_batch_size=reanalyze_batch_size, + reanalyze_partition=reanalyze_partition, + ), + ) + jericho_unizero_config = EasyDict(jericho_unizero_config) + + jericho_unizero_create_config: Dict[str, Any] = dict( + env=dict( + type="jericho", + import_names=["zoo.jericho.envs.jericho_env"], + ), + env_manager=dict(type="base"), + policy=dict( + type="unizero", + import_names=["lzero.policy.unizero"], + ), + ) + jericho_unizero_create_config = EasyDict(jericho_unizero_create_config) + + main_config: EasyDict = jericho_unizero_config + create_config: EasyDict = jericho_unizero_create_config + + from ding.utils import DDPContext + from lzero.config.utils import lz_to_ddp_config + with DDPContext(): + main_config = lz_to_ddp_config(main_config) + main_config.exp_name = ( + f"data_lz/data_unizero_jericho/qwen2.5-0.5B/{env_id}/" + f"uz_qwen_ddp-{gpu_num}gpu_cen{collector_env_num}_rr{replay_ratio}_" + f"ftemp025_{env_id[:8]}_ms{max_steps}_ass-{action_space_size}_" + f"nlayer{num_layers}_embed{embed_dim}_Htrain{num_unroll_steps}-" + f"Hinfer{infer_context_length}_bs{batch_size}_seed{seed}" + ) + from lzero.entry import train_unizero + train_unizero( + [main_config, create_config], + seed=seed, + model_path=main_config.policy.model_path, + max_env_step=max_env_step, + ) + + +if __name__ == "__main__": + """ + Example: + torchrun --nproc_per_node=4 ./zoo/jericho/configs/jericho_unizero_qwen_ddp_config.py + """ + parser = argparse.ArgumentParser(description='Process environment configuration and launch training.') + parser.add_argument( + '--env', + type=str, + help='Identifier of the environment, e.g., detective.z5 or zork1.z5', + default='detective.z5' + ) + parser.add_argument( + '--seed', + type=int, + help='Random seed for reproducibility', + default=0 + ) + parser.add_argument( + '--max_env_step', + type=int, + help='Maximum number of environment steps', + default=int(1e6) + ) + args = parser.parse_args() + + os.environ['TOKENIZERS_PARALLELISM'] = 'false' + main(args.env, args.seed, args.max_env_step) diff --git a/zoo/jericho/envs/jericho_env.py b/zoo/jericho/envs/jericho_env.py index e6ac44a2b..c3d0955d9 100644 --- a/zoo/jericho/envs/jericho_env.py +++ b/zoo/jericho/envs/jericho_env.py @@ -2,8 +2,12 @@ import copy import os import json +import signal as _signal +import time +import multiprocessing as _mp from datetime import datetime from typing import Any, Dict, List, Optional, Union +from collections import OrderedDict import gym import numpy as np @@ -12,7 +16,137 @@ from ding.utils import ENV_REGISTRY, set_pkg_seed, get_rank, get_world_size from ding.envs import BaseEnv, BaseEnvTimestep -from jericho import FrotzEnv + + +def _valid_actions_worker_loop(conn, game_path, seed): + """Persistent worker: receives game states, returns valid actions.""" + os.setpgrp() # new process group so killpg can reach pool workers + try: + from jericho import FrotzEnv as _FrotzEnv + env = _FrotzEnv(game_path, seed) + while True: + try: + state = conn.recv() + if state is None: # shutdown sentinel + break + env.set_state(state) + actions = env.get_valid_actions() + conn.send(actions) + except EOFError: + break + except Exception: + try: + conn.send([]) + except Exception: + break + finally: + conn.close() + +class _ValidActionsWorker: + """ + Manages a long-lived child process that runs get_valid_actions(). + + * First call creates the child (which loads its own FrotzEnv + pool once). + * Subsequent calls just send state / receive actions via Pipe (~0 overhead). + * On timeout the entire process group is SIGKILL'd and a fresh child starts. + """ + def __init__(self, game_path, seed=0): + self.game_path = game_path + self.seed = seed + self._proc: Optional[_mp.Process] = None + self._conn = None + self._start() + + def _start(self): + parent_conn, child_conn = _mp.Pipe() + self._proc = _mp.Process( + target=_valid_actions_worker_loop, + args=(child_conn, self.game_path, self.seed), + ) + self._proc.start() + child_conn.close() # only the worker uses this end + self._conn = parent_conn + + def _kill(self): + if self._proc is not None: + pid = self._proc.pid + # Kill entire process group (child + its pool workers) + try: + os.killpg(pid, _signal.SIGKILL) + except (ProcessLookupError, PermissionError, OSError): + try: + self._proc.kill() + except Exception: + pass + try: + self._proc.join(timeout=5) + except Exception: + pass + if self._conn is not None: + try: + self._conn.close() + except Exception: + pass + self._proc = None + self._conn = None + + def _restart(self): + self._kill() + self._start() + + def close(self): + if self._proc is not None and self._proc.is_alive(): + try: + self._conn.send(None) # graceful shutdown + self._proc.join(timeout=5) + except Exception: + pass + # force-kill if still alive + if self._proc is not None and self._proc.is_alive(): + self._kill() + return + self._kill() # clean up handl + + + def get_valid_actions(self, state, timeout=60): + """ + Send *state* to the worker, wait up to *timeout* seconds. + Returns (actions_list, timed_out). + """ + if self._proc is None or not self._proc.is_alive(): + self._start() + + # Send the state to the worker + try: + self._conn.send(state) + except (BrokenPipeError, OSError): + self._restart() + try: + self._conn.send(state) + except Exception: + return [], True + + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + remaining = max(deadline - time.monotonic(), 0.01) + try: + if self._conn.poll(remaining): + try: + result = self._conn.recv() + return result, False + except Exception: + return [], False + except BaseException: + # SIGALRM or other interruption — keep waiting until deadline + continue + + # Timeout — kill the stuck worker and start a fresh one + logging.warning( + f'[TIMEOUT] get_valid_actions() worker timed out after {timeout}s. ' + f'Killing worker process group and restarting.' + ) + self._restart() + return None, True @ENV_REGISTRY.register('jericho') @@ -49,12 +183,14 @@ class JerichoEnv(BaseEnv): 'max_seq_len': 512, 'remove_stuck_actions': False, 'add_location_and_inventory': False, - # 'for_unizero': False, 'for_unizero': True, 'save_replay': False, 'save_replay_path': None, 'env_type': "zork1", - 'collect_policy_mode': "agent" + 'collect_policy_mode': "agent", + 'use_cache': True, + 'cache_size': 100000, + 'get_valid_actions_timeout': 40, } def __init__(self, cfg: Dict[str, Any]) -> None: @@ -93,6 +229,12 @@ def __init__(self, cfg: Dict[str, Any]) -> None: self.add_location_and_inventory: bool = self.cfg['add_location_and_inventory'] self.for_unizero: bool = self.cfg['for_unizero'] + self.use_cache = self.cfg['use_cache'] + if self.use_cache: + self.cache_size = self.cfg['cache_size'] + self.cache_buffer = OrderedDict() + print(f'[jericho]: use_cache: {self.use_cache}, cache_size={self.cache_size}') + # Initialize the tokenizer once (only in rank 0 process if distributed) if JerichoEnv.tokenizer is None: if self.rank == 0: @@ -103,9 +245,14 @@ def __init__(self, cfg: Dict[str, Any]) -> None: if self.rank != 0: JerichoEnv.tokenizer = AutoTokenizer.from_pretrained(self.cfg['tokenizer_path']) - # Initialize FrotzEnv with the given game. - self._env: FrotzEnv = FrotzEnv(self.game_path, 0) + # Subprocess timeout for Frotz operations (seconds). + # Jericho's C code can hang indefinitely; this is the kill threshold. + self._frotz_timeout: float = float(self.cfg.get('frotz_timeout', 30.0)) + + # Initialize FrotzEnv inside a subprocess for hang protection. + self._env = FrotzWorker(self.game_path, timeout=self._frotz_timeout) self._action_list: Optional[List[str]] = None + self._frotz_halted: bool = False # Set True when Frotz subprocess was killed self.finished: bool = False self._init_flag: bool = False self.episode_return: float = 0.0 @@ -113,7 +260,10 @@ def __init__(self, cfg: Dict[str, Any]) -> None: self._timestep: int = 0 self.episode_history: Optional[List[Dict[str, Any]]] = None self.walkthrough_actions: Optional[List[str]] = None - + + self._get_valid_actions_timeout: bool = False + self._valid_actions_timeout_sec: int = self.cfg.get('get_valid_actions_timeout', 60) + self._valid_actions_worker: Optional[_ValidActionsWorker] = None # Define observation, action, and reward spaces. self.observation_space: gym.spaces.Dict = gym.spaces.Dict() @@ -136,9 +286,49 @@ def prepare_obs(self, obs: str, return_str: bool = False) -> Dict[str, Any]: """ # [PRIORZERO-NEW] Store raw observation text before processing raw_obs_text = obs # Save original text BEFORE any modification - if self._action_list is None: - self._action_list = self._env.get_valid_actions() + if self._valid_actions_worker is None: + self._valid_actions_worker = _ValidActionsWorker( + self.game_path, getattr(self, '_seed', 0) + ) + if self.use_cache: + cache_key = self._env.get_world_state_hash() + if cache_key in self.cache_buffer: + self.cache_buffer.move_to_end(cache_key) + self._action_list = self.cache_buffer[cache_key] + else: + state = self._env.get_state() + actions, timed_out = self._valid_actions_worker.get_valid_actions( + state, timeout=self._valid_actions_timeout_sec + ) + if timed_out: + logging.error( + f'[TIMEOUT] get_valid_actions() timed out after ' + f'{self._valid_actions_timeout_sec}s at timestep ' + f'{self._timestep}! Setting action_list=[] and will end episode.' + ) + self._action_list = [] + self._get_valid_actions_timeout = True + else: + self._action_list = actions if actions is not None else [] + self.cache_buffer[cache_key] = self._action_list + if len(self.cache_buffer) > self.cache_size: + self.cache_buffer.popitem(last=False) + else: + state = self._env.get_state() + actions, timed_out = self._valid_actions_worker.get_valid_actions( + state, timeout=self._valid_actions_timeout_sec + ) + if timed_out: + logging.warning( + f'[TIMEOUT] get_valid_actions() timed out after ' + f'{self._valid_actions_timeout_sec}s at timestep ' + f'{self._timestep}! Setting action_list=[] and will end episode.' + ) + self._action_list = [] + self._get_valid_actions_timeout = True + else: + self._action_list = actions if actions is not None else [] # Filter available actions based on whether stuck actions are removed. if self.remove_stuck_actions: @@ -241,11 +431,13 @@ def reset(self, return_str: bool = False) -> Dict[str, Any]: - (:obj:`Dict[str, Any]`): The processed observation from the environment reset. """ initial_observation, info = self._env.reset() + self._get_valid_actions_timeout = False self.finished = False self._init_flag = True self._action_list = None - self.episode_return = 0.0 + self._frotz_halted = False # Clear halted flag on successful reset + self.episode_return = info['score'] if 'score' in info else 0.0 self._timestep = 0 self.episode_history = [] if self.collect_policy_mode == 'expert': @@ -291,6 +483,8 @@ def close(self) -> None: Close the environment and release any resources. """ self._init_flag = False + if hasattr(self, '_env') and self._env is not None: + self._env.close() def __repr__(self) -> str: """ @@ -317,6 +511,17 @@ def step(self, action: Union[int, np.ndarray, str], return_str: bool = False) -> # Clear previously blocked actions. self.blocked_actions = set() + # If Frotz was previously killed (halted), force done immediately. + if self._frotz_halted: + dummy_obs = self.prepare_obs("[Frotz emulator halted]", return_str) + info = { + 'action_str': 'noop', + 'abnormal': True, + 'frotz_timeout': True, + 'eval_episode_return': self.episode_return, + } + return BaseEnvTimestep(dummy_obs, 0.0, True, info) + # Convert numerical action to string if necessary. if isinstance(action, str): action_str: str = action @@ -343,7 +548,23 @@ def step(self, action: Union[int, np.ndarray, str], return_str: bool = False) -> previous_obs: Optional[str] = self.last_observation if (self.remove_stuck_actions and self.last_observation is not None) else None - observation, reward, done, info = self._env.step(action_str) + try: + observation, reward, done, info = self._env.step(action_str) + except RuntimeError as e: + # FrotzWorker timeout: the Frotz process was killed and respawned. + # Return an abnormal timestep so the caller (BaseEnvManager / Collector) can reset. + logging.warning(f"[JerichoEnv] step() timed out on action '{action_str}': {e}") + self._frotz_halted = True + dummy_obs = self.prepare_obs("[Frotz emulator halted]", return_str) + info = { + 'action_str': action_str, + 'abnormal': True, + 'frotz_timeout': True, + 'eval_episode_return': self.episode_return, + } + return BaseEnvTimestep(dummy_obs, 0.0, True, info) + + info['action_str'] = action_str self._timestep += 1 if not self.for_unizero: @@ -361,6 +582,21 @@ def step(self, action: Union[int, np.ndarray, str], return_str: bool = False) -> self.last_observation = observation processed_obs = self.prepare_obs(observation, return_str) + + # If get_valid_actions timed out during prepare_obs, end the episode. + if self._get_valid_actions_timeout: + done = True + logging.warning( + f'[TIMEOUT] rank {self.rank} get_valid_actions() timed out during step {self._timestep}. ' + f'Ending episode. episode_return: {self.episode_return}' + ) + + # If prepare_obs triggered a timeout (e.g. get_valid_actions hung), force done. + if self._frotz_halted: + info['abnormal'] = True + info['frotz_timeout'] = True + info['eval_episode_return'] = self.episode_return + return BaseEnvTimestep(processed_obs, reward, True, info) if self._timestep >= self.max_steps: done = True @@ -514,13 +750,12 @@ def collect_episode_data(self): if __name__ == '__main__': from easydict import EasyDict - env_type='detective' # zork1, acorncourt, detective, omniquest - # Configuration dictionary for the environment. + env_type = 'zork1' env_cfg = EasyDict( dict( max_steps=400, - game_path="./zoo/jericho/envs/z-machine-games-master/jericho-game-suite/" + f"{env_type}.z5", - max_action_num=10, + game_path="/mnt/afs/niuyazhe/workspace/xiongjyu/LightZero/zoo/jericho/envs/z-machine-games-master/jericho-game-suite/" + f"{env_type}.z5", + max_action_num=200, tokenizer_path="google-bert/bert-base-uncased", max_seq_len=512, remove_stuck_actions=False, @@ -528,10 +763,11 @@ def collect_episode_data(self): for_unizero=False, collector_env_num=1, evaluator_env_num=1, - save_replay=True, + save_replay=False, save_replay_path=None, env_type=env_type, - collect_policy_mode='expert' # random, human, expert + collect_policy_mode='expert', + get_valid_actions_timeout=20, ) ) env = JerichoEnv(env_cfg) diff --git a/zoo/jericho/llm_sft.py b/zoo/jericho/llm_sft.py new file mode 100644 index 000000000..d93a77ecc --- /dev/null +++ b/zoo/jericho/llm_sft.py @@ -0,0 +1,824 @@ +#!/usr/bin/env python3 +import argparse +import json +import os +import random +import re +import shutil +import sys +import tempfile +from collections import deque +from pathlib import Path +from typing import Any, Deque, Dict, Iterable, List, Optional, Sequence, Tuple + +import numpy as np +import torch +import torch.nn.functional as F +from torch.utils.data import Dataset + + +LIGHTZERO_ROOT = Path(__file__).resolve().parents[2] +if str(LIGHTZERO_ROOT) not in sys.path: + sys.path.insert(0, str(LIGHTZERO_ROOT)) + +from jericho.util import unabbreviate as jericho_unabbreviate # noqa: E402 +from zoo.jericho.envs.jericho_env import JerichoEnv # noqa: E402 + + +ENV_PRESETS: Dict[str, Dict[str, int]] = { + "detective.z5": {"max_action_num": 12, "max_steps": 100}, + "omniquest.z5": {"max_action_num": 25, "max_steps": 100}, + "acorncourt.z5": {"max_action_num": 45, "max_steps": 50}, + "zork1.z5": {"max_action_num": 55, "max_steps": 500}, +} + +DEFAULT_ENVS = ["detective.z5", "omniquest.z5", "acorncourt.z5", "zork1.z5"] +DEFAULT_EXPERIMENT_MODES = ["without_valid_actions", "with_valid_actions"] + +SYSTEM_PROMPT = ( + "You are an expert player in a text-based adventure game.\n" + "Your goal is to maximize score by choosing the best next action.\n" + "Always output exactly one line in this format:\n" + "Action: " +) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser() + parser.add_argument("--env_ids", nargs="+", default=DEFAULT_ENVS) + parser.add_argument( + "--experiment_modes", + nargs="+", + default=DEFAULT_EXPERIMENT_MODES, + ) + parser.add_argument( + "--base_model_path", + type=str, + default="/mnt/afs/niuyazhe/workspace/xiongjyu/models/Qwen2.5-3B-Instruct", + ) + parser.add_argument("--output_dir", type=str, default="./outputs/jericho_qwen25_3b_sft") + parser.add_argument("--history_window", type=int, default=10) + parser.add_argument("--collect_episodes_per_env", type=int, default=1) + parser.add_argument("--eval_episodes_per_env", type=int, default=3) + parser.add_argument("--max_seq_len", type=int, default=2048) + parser.add_argument("--max_new_tokens", type=int, default=32) + parser.add_argument("--scoring_batch_size", type=int, default=8) + parser.add_argument("--eval_action_mode", choices=["score", "generate"], default="score") + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--eval_seed", type=int, default=1234) + parser.add_argument("--num_epochs", type=int, default=5) + parser.add_argument("--train_batch_size", type=int, default=4) + parser.add_argument("--learning_rate", type=float, default=1e-5) + parser.add_argument("--weight_decay", type=float, default=0.0) + parser.add_argument("--grad_accum_steps", type=int, default=1) + parser.add_argument("--max_grad_norm", type=float, default=1.0) + parser.add_argument("--warmup_ratio", type=float, default=0.05) + parser.add_argument("--lr_scheduler_type", type=str, default="cosine") + parser.add_argument("--log_every", type=int, default=10) + return parser.parse_args() + + +def set_seed(seed: int) -> None: + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(seed) + + +def reward_to_float(reward: Any) -> float: + if isinstance(reward, np.ndarray): + return float(reward.item()) + if isinstance(reward, torch.Tensor): + return float(reward.item()) + return float(reward) + + +def dump_json(path: str, obj: Any) -> None: + os.makedirs(os.path.dirname(path), exist_ok=True) + with open(path, "w", encoding="utf-8") as f: + json.dump(obj, f, ensure_ascii=False, indent=2) + + +def dump_jsonl(path: str, records: Iterable[Dict[str, Any]]) -> None: + os.makedirs(os.path.dirname(path), exist_ok=True) + with open(path, "w", encoding="utf-8") as f: + for record in records: + f.write(json.dumps(record, ensure_ascii=False) + "\n") + + +def load_jsonl(path: str) -> List[Dict[str, Any]]: + records: List[Dict[str, Any]] = [] + with open(path, "r", encoding="utf-8") as f: + for line in f: + line = line.strip() + if line: + records.append(json.loads(line)) + return records + + +def normalize_action_text(action: str) -> str: + return re.sub(r"\s+", " ", action.strip().lower()) + + +def build_env_cfg(env_id: str, tokenizer_path: str) -> Dict[str, Any]: + if env_id not in ENV_PRESETS: + raise ValueError(f"Unknown env_id={env_id}") + game_path = os.path.join( + "/mnt/afs/niuyazhe/workspace/xiongjyu/LightZero/zoo/jericho/envs/z-machine-games-master/jericho-game-suite", + env_id, + ) + if not os.path.exists(game_path): + raise FileNotFoundError(f"Game file not found: {game_path}") + preset = ENV_PRESETS[env_id] + return { + "max_steps": int(preset["max_steps"]), + "game_path": game_path, + "max_action_num": int(preset["max_action_num"]), + "tokenizer_path": tokenizer_path, + "max_seq_len": 512, + "remove_stuck_actions": False, + "add_location_and_inventory": False, + "for_unizero": False, + "save_replay": False, + "save_replay_path": None, + "env_type": env_id.replace(".z5", ""), + "collect_policy_mode": "expert", + "use_cache": True, + "cache_size": 100000, + } + + +def build_user_prompt( + history: Sequence[Tuple[str, str, float]], + current_obs: str, + valid_actions: Sequence[str], + include_valid_actions: bool, +) -> str: + parts: List[str] = [] + if history: + parts.append("=== GAME HISTORY ===") + for idx, (obs, action, reward) in enumerate(history, start=1): + parts.append(f"Step {idx}:") + parts.append(f"Observation: {obs.strip()}") + parts.append(f"Action: {action.strip()}") + parts.append(f"Reward: {reward:.4f}") + parts.append("") + + parts.append("=== CURRENT OBSERVATION ===") + parts.append(current_obs.strip()) + + if include_valid_actions and valid_actions: + parts.append("") + parts.append("[Valid Actions]") + parts.append("You must choose exactly one action from this list:") + parts.append(", ".join([f"'{action}'" for action in valid_actions])) + + parts.append("") + parts.append("=== INSTRUCTION ===") + parts.append("Output only one line: Action: ") + return "\n".join(parts) + + +def build_chat_prompt(tokenizer: Any, question: str, system_prompt: str = SYSTEM_PROMPT) -> str: + return tokenizer.apply_chat_template( + [ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": question}, + ], + tokenize=False, + add_generation_prompt=True, + ) + + +def extract_action_from_generation(text: str, strict_regex_only: bool = False) -> str: + matches = re.findall(r"Action\s*:\s*([^\n\r]+)", text, flags=re.IGNORECASE) + lines = [line.strip() for line in text.splitlines() if line.strip()] + if matches: + candidate = matches[-1] + elif strict_regex_only: + candidate = lines[0] if lines else "" + else: + candidate = lines[0] if lines else "" + return candidate.strip().strip("`").strip("\"").strip("'").strip() + + +def match_action_in_valid(action: str, valid_actions: Sequence[str]) -> Optional[str]: + if not valid_actions: + return None + valid_map = {normalize_action_text(valid_action): valid_action for valid_action in valid_actions} + return valid_map.get(normalize_action_text(action)) + + +def collect_walkthrough_data_for_env( + env_id: str, + env_cfg: Dict[str, Any], + history_window: int, + collect_episodes_per_env: int, + seed: int, +) -> List[Dict[str, Any]]: + samples: List[Dict[str, Any]] = [] + cfg = dict(env_cfg) + cfg["collect_policy_mode"] = "expert" + env = JerichoEnv(cfg) + + try: + for episode_id in range(collect_episodes_per_env): + env.seed(seed + episode_id, dynamic_seed=False) + obs = env.reset(return_str=True) + history: Deque[Tuple[str, str, float]] = deque(maxlen=history_window) + walkthrough_actions = list(env.walkthrough_actions or []) + + for step_id, action in enumerate(walkthrough_actions): + action = jericho_unabbreviate(str(action)).strip() + current_obs = str(obs.get("raw_obs_text", "")) + valid_actions = [str(item) for item in obs.get("valid_actions", [])] + + samples.append( + { + "env_id": env_id, + "episode_id": episode_id, + "step_id": step_id, + "history": list(history), + "current_obs": current_obs, + "valid_actions": valid_actions, + "target_action": action, + } + ) + + next_obs, reward, done, info = env.step(action, return_str=True) + reward_value = reward_to_float(reward) + executed_action = str(info.get("action_str", action)) + history.append((current_obs, executed_action, reward_value)) + obs = next_obs + if done: + break + finally: + env.close() + + return samples + +def collect_walkthrough_samples(args: argparse.Namespace) -> List[Dict[str, Any]]: + all_samples: List[Dict[str, Any]] = [] + for env_id in args.env_ids: + env_cfg = build_env_cfg(env_id, args.base_model_path) + env_samples = collect_walkthrough_data_for_env( + env_id=env_id, + env_cfg=env_cfg, + history_window=args.history_window, + collect_episodes_per_env=args.collect_episodes_per_env, + seed=args.seed, + ) + all_samples.extend(env_samples) + print(f"[Collect] env={env_id}, samples={len(env_samples)}") + print(f"[Collect] total_samples={len(all_samples)}") + return all_samples + + +def build_train_records( + raw_samples: Sequence[Dict[str, Any]], + include_valid_actions: bool, +) -> List[Dict[str, str]]: + records: List[Dict[str, str]] = [] + for sample in raw_samples: + question = build_user_prompt( + history=sample["history"], + current_obs=str(sample["current_obs"]), + valid_actions=sample["valid_actions"], + include_valid_actions=include_valid_actions, + ) + answer = f"Action: {sample['target_action']}" + records.append({"question": question, "answer": answer}) + return records + + +def prepare_train_jsonl( + raw_samples: Sequence[Dict[str, Any]], + mode_dir: str, + include_valid_actions: bool, +) -> Tuple[str, List[Dict[str, str]]]: + train_jsonl_path = os.path.join(mode_dir, "train.jsonl") + if len(raw_samples) == 0: + raise RuntimeError(f"Missing raw walkthrough samples to build {train_jsonl_path}.") + train_records = build_train_records(raw_samples, include_valid_actions=include_valid_actions) + dump_jsonl(train_jsonl_path, train_records) + return train_jsonl_path, train_records + + +class TrainJsonlDataset(Dataset): + def __init__(self, train_records: Sequence[Dict[str, str]], tokenizer: Any, max_seq_len: int): + self.items: List[Dict[str, List[int]]] = [] + eos_text = tokenizer.eos_token if tokenizer.eos_token is not None else "" + + for record in train_records: + question = str(record["question"]) + answer = str(record["answer"]) + prompt_text = build_chat_prompt(tokenizer, question=question, system_prompt=SYSTEM_PROMPT) + target_text = f"{answer}{eos_text}" + full_text = prompt_text + target_text + + encoded_full = tokenizer( + full_text, + add_special_tokens=False, + truncation=True, + max_length=max_seq_len, + ) + input_ids = encoded_full["input_ids"] + attention_mask = encoded_full["attention_mask"] + labels = [-100] * len(input_ids) + + target_ids = tokenizer(target_text, add_special_tokens=False, truncation=False)["input_ids"] + target_len = min(len(target_ids), len(input_ids)) + labels[-target_len:] = input_ids[-target_len:] + + self.items.append( + { + "input_ids": input_ids, + "attention_mask": attention_mask, + "labels": labels, + } + ) + + def __len__(self) -> int: + return len(self.items) + + def __getitem__(self, idx: int) -> Dict[str, List[int]]: + return self.items[idx] + + +class SFTCollator: + def __init__(self, pad_token_id: int, padding_side: str): + self.pad_token_id = int(pad_token_id) + if padding_side not in {"left", "right"}: + raise ValueError(f"Unsupported padding_side: {padding_side}") + self.padding_side = padding_side + + def __call__(self, features: Sequence[Dict[str, List[int]]]) -> Dict[str, torch.Tensor]: + max_len = max(len(feature["input_ids"]) for feature in features) + input_ids: List[List[int]] = [] + attention_mask: List[List[int]] = [] + labels: List[List[int]] = [] + + for feature in features: + pad_len = max_len - len(feature["input_ids"]) + if self.padding_side == "left": + input_ids.append([self.pad_token_id] * pad_len + feature["input_ids"]) + attention_mask.append([0] * pad_len + feature["attention_mask"]) + labels.append([-100] * pad_len + feature["labels"]) + else: + input_ids.append(feature["input_ids"] + [self.pad_token_id] * pad_len) + attention_mask.append(feature["attention_mask"] + [0] * pad_len) + labels.append(feature["labels"] + [-100] * pad_len) + + return { + "input_ids": torch.tensor(input_ids, dtype=torch.long), + "attention_mask": torch.tensor(attention_mask, dtype=torch.long), + "labels": torch.tensor(labels, dtype=torch.long), + } + + +def load_tokenizer_llm(model_path: str) -> Tuple[Any, Any]: + from transformers import AutoModelForCausalLM, AutoTokenizer + + tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True) + if tokenizer.pad_token is None: + tokenizer.pad_token = tokenizer.eos_token + tokenizer.padding_side = "left" + tokenizer.truncation_side = "left" + + if not torch.cuda.is_available(): + raise RuntimeError("BF16 is required, but CUDA is not available.") + if not torch.cuda.is_bf16_supported(): + raise RuntimeError("BF16 is required, but current CUDA device does not support BF16.") + + model = AutoModelForCausalLM.from_pretrained( + model_path, + trust_remote_code=True, + torch_dtype=torch.bfloat16, + ).to("cuda") + return tokenizer, model + + +def train_sft( + args: argparse.Namespace, + model: Any, + tokenizer: Any, + train_records: Sequence[Dict[str, str]], + work_dir: str, +) -> Dict[str, Any]: + from transformers import Trainer, TrainingArguments + + if len(train_records) == 0: + raise RuntimeError("No training records found.") + + dataset = TrainJsonlDataset(train_records, tokenizer, args.max_seq_len) + trainer_output_dir = tempfile.mkdtemp(prefix="trainer_", dir=work_dir) + original_use_cache = getattr(model.config, "use_cache", None) + if original_use_cache is not None: + model.config.use_cache = False + training_args = TrainingArguments( + output_dir=trainer_output_dir, + overwrite_output_dir=True, + num_train_epochs=args.num_epochs, + per_device_train_batch_size=args.train_batch_size, + gradient_accumulation_steps=args.grad_accum_steps, + learning_rate=args.learning_rate, + weight_decay=args.weight_decay, + max_grad_norm=args.max_grad_norm, + warmup_ratio=args.warmup_ratio, + lr_scheduler_type=args.lr_scheduler_type, + bf16=True, + logging_strategy="steps", + logging_steps=max(1, args.log_every), + save_strategy="no", + report_to=[], + remove_unused_columns=False, + disable_tqdm=False, + ) + trainer = Trainer( + model=model, + args=training_args, + train_dataset=dataset, + data_collator=SFTCollator(tokenizer.pad_token_id, tokenizer.padding_side), + ) + try: + train_result = trainer.train() + trainer.save_state() + metrics = dict(train_result.metrics) + finally: + if original_use_cache is not None: + model.config.use_cache = original_use_cache + shutil.rmtree(trainer_output_dir, ignore_errors=True) + return metrics + + +class JerichoLLMAgent: + def __init__( + self, + model: Any, + tokenizer: Any, + device: torch.device, + max_seq_len: int, + max_new_tokens: int, + scoring_batch_size: int, + action_mode: str, + include_valid_actions: bool, + ): + self.model = model + self.tokenizer = tokenizer + self.device = device + self.max_seq_len = max_seq_len + self.max_new_tokens = max_new_tokens + self.scoring_batch_size = scoring_batch_size + self.action_mode = action_mode + self.include_valid_actions = include_valid_actions + self.eos_text = tokenizer.eos_token if tokenizer.eos_token is not None else "" + + @torch.no_grad() + def _score_actions(self, chat_prompt: str, valid_actions: Sequence[str]) -> Dict[str, float]: + if not valid_actions: + return {"go": 0.0} + + scores: Dict[str, float] = {} + completions = [f"Action: {action}{self.eos_text}" for action in valid_actions] + for start in range(0, len(valid_actions), self.scoring_batch_size): + end = min(start + self.scoring_batch_size, len(valid_actions)) + batch_actions = list(valid_actions[start:end]) + batch_completions = completions[start:end] + full_texts = [chat_prompt + completion for completion in batch_completions] + + enc = self.tokenizer( + full_texts, + return_tensors="pt", + padding=True, + truncation=True, + max_length=self.max_seq_len, + add_special_tokens=False, + ) + input_ids = enc["input_ids"].to(self.device) + attention_mask = enc["attention_mask"].to(self.device) + + labels = torch.full_like(input_ids, -100) + for idx, completion in enumerate(batch_completions): + completion_ids = self.tokenizer(completion, add_special_tokens=False, truncation=False)["input_ids"] + nonpad_pos = torch.nonzero(attention_mask[idx], as_tuple=False).squeeze(-1) + if nonpad_pos.numel() == 0: + continue + seq_end = int(nonpad_pos[-1].item()) + 1 + target_len = min(len(completion_ids), seq_end) + labels[idx, seq_end - target_len : seq_end] = input_ids[idx, seq_end - target_len : seq_end] + + outputs = self.model(input_ids=input_ids, attention_mask=attention_mask) + logits = outputs.logits[:, :-1, :] + shifted_ids = input_ids[:, 1:] + shifted_labels = labels[:, 1:] + valid_mask = shifted_labels.ne(-100) + token_logprobs = F.log_softmax(logits, dim=-1).gather( + dim=-1, + index=shifted_ids.unsqueeze(-1), + ).squeeze(-1) + denom = valid_mask.sum(dim=1).clamp(min=1) + score_tensor = (token_logprobs * valid_mask).sum(dim=1) / denom + + for action, score in zip(batch_actions, score_tensor.detach().cpu().tolist()): + scores[action] = float(score) + return scores + + @torch.no_grad() + def _generate_action(self, chat_prompt: str) -> str: + enc = self.tokenizer( + chat_prompt, + return_tensors="pt", + truncation=True, + max_length=self.max_seq_len, + add_special_tokens=False, + ) + input_ids = enc["input_ids"].to(self.device) + attention_mask = enc["attention_mask"].to(self.device) + out = self.model.generate( + input_ids=input_ids, + attention_mask=attention_mask, + max_new_tokens=self.max_new_tokens, + do_sample=False, + eos_token_id=self.tokenizer.eos_token_id, + pad_token_id=self.tokenizer.pad_token_id, + ) + gen_ids = out[0, input_ids.size(1) :] + return self.tokenizer.decode(gen_ids, skip_special_tokens=True) + + def select_action( + self, + history: Sequence[Tuple[str, str, float]], + current_obs: str, + valid_actions: Sequence[str], + ) -> Tuple[str, str, str]: + question = build_user_prompt( + history=history, + current_obs=current_obs, + valid_actions=valid_actions, + include_valid_actions=self.include_valid_actions, + ) + chat_prompt = build_chat_prompt(self.tokenizer, question=question) + + if not self.include_valid_actions: + raw_generation = self._generate_action(chat_prompt) + action_str = extract_action_from_generation(raw_generation, strict_regex_only=True) + return action_str, "generate_no_valid_direct", raw_generation + + if not valid_actions: + return "go", "fallback_no_valid_actions", "" + + if self.action_mode == "score": + scores = self._score_actions(chat_prompt, valid_actions) + return max(scores.items(), key=lambda item: item[1])[0], "score", "" + + raw_generation = self._generate_action(chat_prompt) + predicted_action = extract_action_from_generation(raw_generation) + mapped_action = match_action_in_valid(predicted_action, valid_actions) + if mapped_action is None: + mapped_action = match_action_in_valid(jericho_unabbreviate(predicted_action), valid_actions) + if mapped_action is not None: + return mapped_action, "generate", raw_generation + + scores = self._score_actions(chat_prompt, valid_actions) + return max(scores.items(), key=lambda item: item[1])[0], "generate_fallback_score", raw_generation + + +def summarize_eval( + stage_name: str, + include_valid_actions: bool, + effective_action_policy: str, + per_env_scores: Dict[str, List[float]], + per_env_returns: Dict[str, List[float]], +) -> Dict[str, Any]: + overall_scores = [score for values in per_env_scores.values() for score in values] + overall_returns = [ret for values in per_env_returns.values() for ret in values] + summary: Dict[str, Any] = { + "stage": stage_name, + "include_valid_actions": include_valid_actions, + "effective_action_policy": effective_action_policy, + "overall": { + "num_episodes": len(overall_scores), + "score_mean": float(np.mean(overall_scores)) if overall_scores else 0.0, + "score_std": float(np.std(overall_scores)) if overall_scores else 0.0, + "return_mean": float(np.mean(overall_returns)) if overall_returns else 0.0, + "return_std": float(np.std(overall_returns)) if overall_returns else 0.0, + }, + "per_env": {}, + } + for env_id in per_env_scores: + env_scores = per_env_scores[env_id] + env_returns = per_env_returns[env_id] + summary["per_env"][env_id] = { + "num_episodes": len(env_scores), + "score_mean": float(np.mean(env_scores)) if env_scores else 0.0, + "score_std": float(np.std(env_scores)) if env_scores else 0.0, + "return_mean": float(np.mean(env_returns)) if env_returns else 0.0, + "return_std": float(np.std(env_returns)) if env_returns else 0.0, + } + return summary + + +def evaluate_model( + args: argparse.Namespace, + model: Any, + tokenizer: Any, + stage_name: str, + stage_dir: str, + include_valid_actions: bool, +) -> Dict[str, Any]: + os.makedirs(stage_dir, exist_ok=True) + agent = JerichoLLMAgent( + model=model, + tokenizer=tokenizer, + device=model.device, + max_seq_len=args.max_seq_len, + max_new_tokens=args.max_new_tokens, + scoring_batch_size=args.scoring_batch_size, + action_mode=args.eval_action_mode, + include_valid_actions=include_valid_actions, + ) + + effective_action_policy = "generate_direct" if not include_valid_actions else args.eval_action_mode + episode_records: List[Dict[str, Any]] = [] + per_env_scores: Dict[str, List[float]] = {env_id: [] for env_id in args.env_ids} + per_env_returns: Dict[str, List[float]] = {env_id: [] for env_id in args.env_ids} + + for env_id in args.env_ids: + env_cfg = build_env_cfg(env_id, args.base_model_path) + env_cfg["collect_policy_mode"] = "agent" + env = JerichoEnv(env_cfg) + + try: + for episode_id in range(args.eval_episodes_per_env): + env.seed(args.eval_seed + episode_id, dynamic_seed=False) + obs = env.reset(return_str=True) + history: Deque[Tuple[str, str, float]] = deque(maxlen=args.history_window) + trajectory: List[Dict[str, Any]] = [] + last_info: Dict[str, Any] = {} + step_count = 0 + + while True: + current_obs = str(obs.get("raw_obs_text", "")) + valid_actions = [str(action) for action in obs.get("valid_actions", [])] + selected_action, selection_method, raw_generation = agent.select_action( + history=list(history), + current_obs=current_obs, + valid_actions=valid_actions, + ) + + next_obs, reward, done, info = env.step(selected_action, return_str=True) + reward_value = reward_to_float(reward) + executed_action = str(info.get("action_str", selected_action)) + + trajectory.append( + { + "step_id": step_count, + "observation": current_obs, + "selected_action": selected_action, + "executed_action": executed_action, + "reward": reward_value, + "score": float(info.get("score", 0.0)), + "done": bool(done), + "selection_method": selection_method, + "raw_generation": raw_generation, + } + ) + + history.append((current_obs, executed_action, reward_value)) + obs = next_obs + last_info = info + step_count += 1 + if done: + break + + episode_return = float(last_info.get("eval_episode_return", env.episode_return)) + episode_score = float(last_info.get("score", 0.0)) + per_env_scores[env_id].append(episode_score) + per_env_returns[env_id].append(episode_return) + episode_records.append( + { + "stage": stage_name, + "env_id": env_id, + "episode_id": episode_id, + "seed": args.eval_seed + episode_id, + "score": episode_score, + "episode_return": episode_return, + "steps": step_count, + "trajectory": trajectory, + } + ) + finally: + env.close() + + episode_path = os.path.join(stage_dir, "eval_episode.jsonl") + dump_jsonl(episode_path, episode_records) + + summary = summarize_eval( + stage_name=stage_name, + include_valid_actions=include_valid_actions, + effective_action_policy=effective_action_policy, + per_env_scores=per_env_scores, + per_env_returns=per_env_returns, + ) + summary["eval_episode_path"] = episode_path + dump_json(os.path.join(stage_dir, "eval_return.json"), summary) + return summary + + +def run_mode_experiment( + args: argparse.Namespace, + mode_name: str, + include_valid_actions: bool, + raw_samples: Sequence[Dict[str, Any]], +) -> Dict[str, Any]: + mode_dir = os.path.join(args.output_dir, mode_name) + os.makedirs(mode_dir, exist_ok=True) + + train_jsonl_path, train_records = prepare_train_jsonl( + raw_samples=raw_samples, + mode_dir=mode_dir, + include_valid_actions=include_valid_actions, + ) + print(f"[Mode:{mode_name}] train_jsonl={train_jsonl_path}, records={len(train_records)}") + + tokenizer, model = load_tokenizer_llm(args.base_model_path) + try: + model.eval() + pre_summary = evaluate_model( + args=args, + model=model, + tokenizer=tokenizer, + stage_name="pre_sft", + stage_dir=os.path.join(mode_dir, "pre_sft"), + include_valid_actions=include_valid_actions, + ) + + train_metrics = train_sft( + args=args, + model=model, + tokenizer=tokenizer, + train_records=train_records, + work_dir=mode_dir, + ) + + model.eval() + post_summary = evaluate_model( + args=args, + model=model, + tokenizer=tokenizer, + stage_name="post_sft", + stage_dir=os.path.join(mode_dir, "post_sft"), + include_valid_actions=include_valid_actions, + ) + finally: + del model + torch.cuda.empty_cache() + + return { + "mode_name": mode_name, + "include_valid_actions": include_valid_actions, + "train_jsonl": train_jsonl_path, + "pre_sft": pre_summary, + "post_sft": post_summary, + "train_metrics": train_metrics, + } + + +def main() -> None: + args = parse_args() + set_seed(args.seed) + + os.makedirs(args.output_dir, exist_ok=True) + print(f"[Config] output_dir={args.output_dir}") + print(f"[Config] env_ids={args.env_ids}") + print(f"[Config] experiment_modes={args.experiment_modes}") + print(f"[Config] history_window={args.history_window}") + + raw_samples = collect_walkthrough_samples(args) + + mode_reports: List[Dict[str, Any]] = [] + for mode_name in args.experiment_modes: + include_valid_actions = mode_name == "with_valid_actions" + mode_report = run_mode_experiment( + args=args, + mode_name=mode_name, + include_valid_actions=include_valid_actions, + raw_samples=raw_samples, + ) + mode_reports.append(mode_report) + + print("[Summary]") + for report in mode_reports: + pre_score = report["pre_sft"]["overall"]["score_mean"] + post_score = report["post_sft"]["overall"]["score_mean"] + pre_return = report["pre_sft"]["overall"]["return_mean"] + post_return = report["post_sft"]["overall"]["return_mean"] + print( + f" - {report['mode_name']}: " + f"score_mean {pre_score:.3f} -> {post_score:.3f}, " + f"return_mean {pre_return:.3f} -> {post_return:.3f}" + ) + +if __name__ == "__main__": + main() diff --git a/zoo/jericho/priorzero/README.md b/zoo/jericho/priorzero/README.md index 7c5b7ddd6..83dfdaf55 100644 --- a/zoo/jericho/priorzero/README.md +++ b/zoo/jericho/priorzero/README.md @@ -1,599 +1,213 @@ -# PriorZero: LLM-Guided World Model Planning +# PriorZero说明 -**PriorZero** combines large language models (LLMs) with world model-based planning (UniZero) for efficient decision-making in complex text-based environments. +本文档面向 `zoo/jericho/priorzero` 分支代码,重点补充: +1. 主要文件说明; +2. 主实验启动命令与修改点; +3. 当前实验结论; +4. 后续可能改进方向。 -## 🎯 Core Idea - -**Decouple Policy and World Model:** -- **LLM Policy**: Provides high-quality action priors using language understanding and world knowledge -- **World Model (UniZero)**: Performs efficient multi-step planning in latent space via MCTS - -**Training Loop:** -1. **Collect**: LLM generates action rankings → MCTS search refines them → Execute best action -2. **Store**: Save MCTS visit distributions (for SFT) and environment rewards (for RFT) -3. **Train**: - - World Model: Standard UniZero losses (value, policy, reward, latent) - - LLM: Supervised Fine-Tuning (SFT) on MCTS policies + Reinforcement Fine-Tuning (RFT) on env rewards - -## 📁 File Structure - -``` -priorzero/ -├── priorzero_entry.py # Main async training loop (stable, tested) -├── priorzero_orz_complete.py # ORZ integration version (experimental) -├── priorzero_config.py # Complete configuration with presets -├── priorzero_policy.py # Dual-model policy (World Model + LLM) -├── priorzero_collector.py # Async data collection with vLLM -├── game_segment_priorzero.py # Enhanced GameSegment with MCTS policies & raw text -├── ensure_local_lightzero.py # Import path management -└── README.md # This file -``` - -## 🔀 Two Training Entry Points - -PriorZero provides two training entry points with different LLM training strategies: - -### 1. `priorzero_entry.py` - Standard PriorZero (Stable ✅) - -**Status**: Production-ready, tested, can run for extended periods - -**LLM Training Strategy**: -- **Built-in SFT + RFT** implemented directly in `priorzero_policy.py` -- Uses micro-batching with gradient accumulation (memory efficient) -- Simple and straightforward implementation -- Fully integrated with UniZero training loop - -**Key Features**: -- Single-process async training -- vLLM for inference only (action prior generation) -- LLM training via standard PyTorch optimizer -- ~580 lines of clean, maintainable code - -**When to use**: -- ✅ Standard PriorZero experiments -- ✅ Quick prototyping and debugging -- ✅ Single GPU training -- ✅ When you want simple, stable training - -**Usage**: -```bash -# Quick test -python priorzero_entry.py --quick_test --env_id zork1.z5 --seed 0 - -# Full training -python priorzero_entry.py --env_id zork1.z5 --seed 0 --max_iter 100000 -``` - -### 2. `priorzero_orz_complete.py` - ORZ Integration (Experimental ⚠️) - -**Status**: Newly implemented, requires testing, not yet verified - -**LLM Training Strategy**: -- **ORZ RayPPOTrainer** for distributed PPO-based LLM fine-tuning -- Leverages OpenAI's ORZ (Open Reasoner Zero) framework -- More sophisticated RL training with actor-critic architecture -- Distributed training with Ray - -**Key Features**: -- Hybrid training: UniZero world model + ORZ PPO for LLM -- Ray-based distributed execution -- Custom reward function for Jericho text adventures -- Separate training frequencies for world model vs LLM -- ~960 lines with complete ORZ integration - -**Key Differences from Standard Entry**: -1. **LLM Training**: Uses ORZ's `RayPPOTrainer` instead of built-in SFT/RFT -2. **Reward Signal**: Custom `JerichoRewardTrainer` for text adventure rewards -3. **Distribution**: Ray-based parallel training -4. **Complexity**: More sophisticated but requires ORZ dependency -5. **Training Loop**: Separate update frequencies for WM and LLM - -**When to use**: -- ⚠️ Advanced RL research with PPO-based LLM training -- ⚠️ When you have ORZ framework available -- ⚠️ Distributed training across multiple GPUs/nodes -- ⚠️ When you want more sophisticated reward modeling - -**Requirements**: -```bash -# Additional dependencies -pip install ray # For distributed execution -cd /path/to/Open-Reasoner-Zero && pip install -e . -``` - -**Usage**: -```bash -# Debug mode -DEBUG_MODE=True python priorzero_orz_complete.py - -# Full training (requires ORZ setup) -python priorzero_orz_complete.py --env_id zork1.z5 --seed 0 -``` - -### Comparison Table - -| Feature | `priorzero_entry.py` | `priorzero_orz_complete.py` | -|---------|---------------------|----------------------------| -| **Status** | ✅ Stable, Tested | ⚠️ Experimental, Needs Testing | -| **Lines of Code** | ~580 | ~960 | -| **LLM Training** | Built-in SFT+RFT | ORZ RayPPOTrainer (PPO) | -| **Dependencies** | Basic (vLLM, torch) | Advanced (ORZ, Ray) | -| **Training Mode** | Single-process async | Distributed (Ray) | -| **Memory Efficiency** | Micro-batching | Ray workers | -| **Reward Modeling** | Simple env rewards | Custom reward functions | -| **Setup Complexity** | Low | Medium-High | -| **Debugging** | Easy | More complex | -| **Performance** | Not fully verified | Unknown (needs testing) | -| **Recommended For** | Most users | Advanced research | - -### Which One Should You Use? - -**Start with `priorzero_entry.py` if:** -- You're new to PriorZero -- You want stable, tested code -- You're doing standard MCTS + LLM experiments -- You have limited GPU resources -- You want simple debugging - -**Try `priorzero_orz_complete.py` if:** -- You have ORZ framework set up -- You want distributed training -- You need custom reward modeling -- You're doing advanced RL research -- You're willing to debug experimental code - -**Note**: The standard entry (`priorzero_entry.py`) has been tested and can run for extended periods. The ORZ version is newly implemented and requires thorough testing before production use. - - -## 🚀 Quick Start - -### 1. Installation - -**Basic Installation** (for `priorzero_entry.py`): -```bash -# Core dependencies -pip install torch transformers vllm peft -pip install ding-engine tensorboardX loguru easydict jericho - -# LightZero (local development mode) -cd /path/to/LightZero && pip install -e . -``` - -**Advanced Installation** (for `priorzero_orz_complete.py`): -```bash -# Basic dependencies (same as above) -pip install torch transformers vllm peft -pip install ding-engine tensorboardX loguru easydict jericho - -# Additional ORZ dependencies -pip install ray # For distributed training -cd /path/to/Open-Reasoner-Zero && pip install -e . - -# LightZero -cd /path/to/LightZero && pip install -e . -``` - -### 2. Quick Test Run - -**Standard PriorZero** (recommended for most users): -```bash -cd /mnt/nfs/zhangjinouwen/puyuan/LightZero/zoo/jericho/priorzero - -# Quick test (reduced resources, 2 envs, 10 iters) -python priorzero_entry.py --quick_test --env_id zork1.z5 --seed 0 - -# Full training (default: 4 envs, 100k iters) -python priorzero_entry.py --env_id zork1.z5 --seed 0 --max_iter 100000 -``` - -**ORZ Integration** (experimental, requires ORZ setup): -```bash -cd /mnt/nfs/zhangjinouwen/puyuan/LightZero/zoo/jericho/priorzero - -# Debug mode (minimal resources) -DEBUG_MODE=True python priorzero_orz_complete.py - -# Full training with ORZ -python priorzero_orz_complete.py --env_id zork1.z5 --seed 0 -``` - -### 3. Test Individual Components - -```bash -# Test configuration -python priorzero_config.py - -# Test game segment -python game_segment_priorzero.py - -# Test buffer -python ../../../lzero/mcts/buffer/game_buffer_priorzero.py -``` - -## 🔧 Configuration - -### Preset Configurations - -```python -# 1. Standard PriorZero (World Model + LLM with SFT + RFT) -from priorzero_config import get_priorzero_config -main_cfg, create_cfg = get_priorzero_config(env_id='zork1.z5', seed=0) - -# 2. Quick Test (reduced resources) -from priorzero_config import get_priorzero_config_for_quick_test -test_cfg, create_cfg = get_priorzero_config_for_quick_test(env_id='zork1.z5', seed=0) - -# 3. Pure UniZero (no LLM) -from priorzero_config import get_config_pure_unizero -cfg, _ = get_config_pure_unizero() - -# 4. LLM with only SFT (no RFT) -from priorzero_config import get_config_llm_only_sft -cfg, _ = get_config_llm_only_sft() - -# 5. LLM with LoRA (memory efficient) -from priorzero_config import get_config_with_lora -cfg, _ = get_config_with_lora() -``` - -## 📊 Key Features - -### 1. Dual-Model Training - -**World Model (UniZero)**: -- Transformer-based world model in latent space -- Predicts: next latent state, reward, value, policy -- Trained with standard UniZero losses (full batch size) -- **Training frequency**: Every iteration (standard RL loop) - -**LLM Policy** - Two Implementations: - -#### Standard Entry (`priorzero_entry.py`): -- Pre-trained LLM (default: Qwen2.5-0.5B-Instruct) -- Fine-tuned with: - - **SFT**: Supervised by MCTS visit distributions - - **RFT**: Reinforced by environment rewards (REINFORCE) -- **Gradient Accumulation**: Micro-batching to avoid OOM -- **Training frequency**: Every iteration (joint optimization with world model) -- Optional LoRA for parameter-efficient fine-tuning - -#### ORZ Entry (`priorzero_orz_complete.py`): -- Pre-trained LLM (configurable) -- Fine-tuned with: - - **ORZ PPO**: Proximal Policy Optimization via RayPPOTrainer - - **Custom Rewards**: JerichoRewardTrainer for text adventure scoring - - **Actor-Critic**: Separate value network for advantage estimation -- **Ray Distribution**: Parallel workers for distributed training -- **Training frequency**: Configurable (default: every N world model updates) -- Support for LoRA and other PEFT methods - -### 2. Memory-Efficient Training (OOM Fix) - -**Micro-Batching with Gradient Accumulation** (Standard Entry): -```python -llm_policy_cfg = dict( - llm_micro_batch_size=4, # Small batch per forward pass - llm_gradient_accumulation_steps=8, # Accumulate over 8 steps - # Effective batch size = 4 * 8 = 32 -) -``` +--- -**How it works**: -- LLM training processes data in small chunks (2-4 samples) -- Gradients accumulate across micro-batches -- Single optimizer step applies accumulated gradients -- World model still trains with full batches (no slowdown) -- Automatic memory cleanup: `torch.cuda.empty_cache()` after each micro-batch - -**Ray Workers** (ORZ Entry): -- Distributed across multiple Ray actors -- Each worker handles subset of data -- Automatic load balancing -- More scalable for large-scale training - -**Tuning guidelines**: -- **If OOM**: Reduce `llm_micro_batch_size` to 1 or 2 -- **If have more memory**: Increase to 8 or 16 -- Effective batch = `llm_micro_batch_size * llm_gradient_accumulation_steps` - -### 3. LLM-Guided MCTS - -1. LLM generates ranked actions: `[action_1, action_2, ...]` -2. Convert to policy prior: `prior_policy = softmax(weights)` -3. Inject into MCTS root node (replace policy logits) -4. MCTS search refines the policy (25 simulations) -5. Select best action based on visit counts - -### 4. Async Data Collection - -- **vLLM Engine**: Efficient batch inference (V1 API) -- **Error Handling**: Auto-retry (max 3 attempts) with backoff -- **Timeout Control**: 30s default per batch -- **History Buffer**: Sliding window (5 recent transitions) -- **Text Observation**: Properly extracts and stores raw text in `raw_obs_segment` - -### 5. Enhanced Game Buffer - -**PriorZeroGameBuffer** (optimized): -- Overrides `_sample_orig_data()` to cache game segments -- Avoids double sampling (~50% faster) -- Returns `[current_batch, target_batch, game_segments]` -- Minimal memory overhead (uses references, not copies) - -## 🎛️ Key Hyperparameters - -### World Model -```python -world_model_cfg = dict( - num_layers=2, # Transformer layers (reduced for speed) - num_heads=8, # Attention heads - embed_dim=512, # Embedding dimension - context_length=8, # Number of past transitions (2 * infer_context_length) - num_unroll_steps=10, # Unroll steps for training - game_segment_length=50, # Segment length (reduced for quick test) -) -``` +## 1. 主要文件说明 + +### 1.1 `src/priorzero_config.py` +该文件负责**统一管理实验配置**,是 PriorZero 训练流程的入口配置源。主要职责: +- 定义可选 LLM 模型预设(`MODEL_CONFIGS`); +- 定义 `PriorZeroLLMConfig`(LLM/RFT 训练相关核心参数); +- 通过 `get_priorzero_config(...)` 组装环境、策略、采集、评估的总配置。 + +#### `llm_config`(`PriorZeroLLMConfig`)参数逐项说明 +> 下述参数是当前 PriorZero LLM/WM 联合训练最关键的调参入口。建议每次实验先固定大结构(训练模式、模型规模),再小步调整损失与采样超参数。 + +##### A. 基础开关与模型路径 +- `model_name_or_path`:LLM 的 HuggingFace/本地模型路径。 +- `enable_rft`:是否开启 LLM 的 RFT(强化微调)训练。 +- `enable_world_model`:是否开启 World Model(WM)训练。 + +##### B. LLM 训练方式(`train_mode_dict`) +- `mode`:LLM 训练模式,`full` 为全参微调,`lora` 为 LoRA 微调。 +- `lora_r`:LoRA 低秩分解 rank。 +- `lora_alpha`:LoRA 缩放系数。 +- `lora_dropout`:LoRA 路径 dropout 比例。 +- `lora_bias`:LoRA 中 bias 训练策略(`none/all/lora_only`)。 +- `lora_target_modules`:应用 LoRA 的模块列表(如 `q_proj/k_proj/...`)。 + +##### C. 交替训练调度(`train_schedule`) +- `alternate`:是否采用 WM/LLM 严格交替训练。 +- `wm_update_iters`:在交替模式下,每轮 WM 连续更新步数。 +- `llm_update_iters`:在交替模式下,每轮 LLM 连续更新步数。 +- `start_phase`:交替训练起始阶段(`wm` 或 `llm`)。 +- `llm_collect_mode`:LLM 训练阶段的数据采集策略(`wm_collect/wm_llm_collect/no_collect`)。 + +##### D. MCTS 根节点先验融合 +- `llm_prior_temperature`:LLM 先验分布温度(温度越高越平滑)。 +- `mcts_root_logits_dict.mode`:根节点 logits 融合模式(仅 LLM、仅 WM、或二者融合)。 +- `mcts_root_logits_dict.plus_method`:融合权重策略(`fixed` 或 `adaptive`)。 +- `mcts_root_logits_dict.wm_weight`:`fixed` 时 WM 的固定权重。 +- `mcts_root_logits_dict.llm_max_weight`:`adaptive` 时 LLM 最大权重。 +- `mcts_root_logits_dict.llm_min_weight`:`adaptive` 时 LLM 最小权重。 +- `mcts_root_logits_dict.max_envsteps`:`adaptive` 权重衰减参考的总环境步数。 + +##### E. 评估策略(`eval_dict`) +- `eval_dict.world_model`:启用“仅 WM”评估。 +- `eval_dict.world_model_llm_prior`:启用“WM + LLM 先验”评估。 +- `eval_dict.llm_prior`:启用“仅 LLM 先验”评估。 +- `eval_dict.wm_eval_freq`:WM 评估频率。 +- `eval_dict.llm_eval_freq`:LLM 评估频率。 + +##### F. Prompt / 序列相关 +- `attn_implementation`:注意力实现方式(如 `flash_attention_2`)。 +- `history_length`:输入历史轨迹长度。 +- `use_cot`:是否启用 CoT 推理。 +- `cot_weight`:CoT 前缀 token 在损失中的权重。 +- `user_prompt_dict.history_with_reward`:prompt 中是否拼接历史 reward。 +- `user_prompt_dict.observation_with_valid_actions`:prompt 中是否拼接当前合法动作。 +- `prompt_max_len`:输入最大 token 长度。 +- `generate_max_len`:生成最大 token 长度。 +- `bf16`:是否使用 bfloat16。 + +##### G. vLLM 推理与采样 +- `enable_vllm`:是否启用 vLLM 引擎。 +- `enable_prefix_caching`:是否启用前缀缓存。 +- `use_cuda_ipc`:是否使用 CUDA IPC。 +- `enable_vllm_is_correction`:是否启用 vLLM 截断修正逻辑。 +- `vllm_is_truncated_threshold`:vLLM 截断判定阈值区间。 +- `use_mispo`:是否启用 MISPO 相关策略。 +- `mispo_token_truncated_threshold`:MISPO token 级截断阈值。 +- `mispo_traj_truncated_threshold`:MISPO 轨迹级截断阈值。 +- `vllm_sync_backend`:vLLM 参数同步后端(如 `nccl`)。 +- `vllm_tensor_parallel_size`:单个 vLLM engine 的张量并行卡数。 +- `gpu_memory_utilization`:vLLM 可用显存占比。 +- `vllm_enable_sleep`:空闲时是否允许 vLLM 休眠。 +- `temperature`:采样温度。 +- `top_p`:核采样阈值。 +- `seed`:随机种子。 +- `reduction`:损失聚合方式(如 `mean`)。 + +##### H. DeepSpeed / 梯度控制 +- `deepspeed_enable_sleep`:DeepSpeed 相关休眠优化开关。 +- `zero_stage`:DeepSpeed ZeRO stage。 +- `gradient_checkpointing`:是否启用梯度检查点。 +- `gradient_checkpointing_use_reentrant`:梯度检查点 reentrant 配置。 +- `max_norm`:梯度裁剪阈值。 +- `ds_tensor_parallel_size`:DeepSpeed 张量并行规模。 + +##### I. 批大小与数据新鲜度 +- `train_batch_size`:全局训练 batch size。 +- `micro_train_batch_size`:单次前向/反向 micro batch size。 +- `max_rollout_staleness`:rollout 到训练的最大“离线陈旧度”。 + +##### J. 优化器与学习率 +- `learning_rate`:学习率。 +- `adam_betas`:Adam beta 系数。 +- `weight_decay`:权重衰减。 +- `lr_scheduler`:学习率调度器类型。 +- `lr_warmup_ratio`:warmup 占总步数比例。 +- `max_steps`:LLM 训练总步数上限。 + +##### K. 策略优化目标 +- `policy_loss_type`:策略损失类型(`ppo/gspo`)。 +- `reward_func.format_reward`:是否启用格式奖励。 +- `reward_func.format_param.format_weight`:格式奖励权重(adv 权重约为 `1-format_weight`)。 +- `advantage_type`:advantage 定义/归一化方式。 +- `eps_clip_low_high`:PPO clip 范围。 +- `rft_kl_coef`:RFT KL 正则系数。 +- `entropy_loss_coef`:熵奖励系数。 +- `kl_estimator`:KL 估计方法。 + +##### L. 保存与数值稳定 +- `llm_save_freq`:LLM checkpoint 保存频率。 +- `save_path`:模型保存路径(通常被 `exp_name` 目录覆盖)。 +- `value_norm_cfg.enable_stability_optimizer`:是否启用稳定性优化器。 +- `value_norm_cfg.value_norm_init_momentum`:value norm 初期动量。 +- `value_norm_cfg.value_norm_final_momentum`:value norm 后期动量。 +- `value_norm_cfg.value_norm_warmup_steps`:动量从初期到后期的过渡步数。 +- `value_norm_cfg.value_norm_clip_percentile`:value clipping 分位点。 +- `value_norm_cfg.value_norm_clip_method`:value clipping 方法。 +- `value_norm_cfg.value_norm_history_size`:value norm 历史缓存长度。 -### LLM Policy -```python -llm_policy_cfg = dict( - pretrain_llm_path="Qwen/Qwen2.5-0.5B-Instruct", - llm_learning_rate=1e-6, - llm_loss_weight=0.5, # Weight of SFT loss - rft_loss_weight=0.3, # Weight of RFT loss - - # Memory optimization - llm_micro_batch_size=4, # Micro-batch size (2 for quick test) - llm_gradient_accumulation_steps=8, # Accumulation steps (4 for quick test) - - # Prompting - prompt_max_len=2048, # Max prompt length (1024 for quick test) - generate_max_len=256, # Max generation length (128 for quick test) - history_length=5, # Context window (3 for quick test) - use_cot=True, # Chain-of-thought prompting - - # Training strategy - sft_target='mcts_policy', # Supervised by MCTS visit distributions - enable_rft=True, # Enable RFT with env rewards - - # vLLM - gpu_memory_utilization=0.3, # GPU memory fraction for vLLM -) -``` +--- -### MCTS -```python -mcts_cfg = dict( - num_simulations=25, # MCTS simulations per step (10 for quick test) - root_dirichlet_alpha=0.3, # Exploration noise - root_noise_weight=0.25, # Noise weight - pb_c_base=19652, # UCB constants - pb_c_init=1.25, -) -``` +### 1.2 `src/priorzero_entry_sync.py` +该文件是**单进程/主控同步训练入口**,核心流程包括: +- 初始化环境、policy、collector、evaluator、replay buffer; +- 构建 vLLM、PolicyModel、ReferenceModel 与 LLM trainer; +- 执行“数据收集 → WM 训练 → LLM 训练 → 评估”的循环; +- 在交替模式下按照 `train_schedule` 在 `wm/llm` 两阶段切换。 -### Training -```python -training_cfg = dict( - batch_size=64, # World model batch size (32 for quick test) - update_per_collect=10, # Updates per collection cycle (5 for quick test) - max_env_step=1e6, # Max environment steps - eval_freq=500, # Evaluation frequency - - # Replay buffer - replay_buffer_size=10000, - use_priority=True, # Prioritized experience replay - priority_prob_alpha=0.6, - priority_prob_beta=0.4, -) -``` +适用场景:快速调试、单节点控制逻辑验证、定位数据流问题。 -## 📈 Expected Results +### 1.3 `src/priorzero_entry_sync_ddp.py` +该文件是**DDP 多卡同步训练入口**,在 `priorzero_entry_sync.py` 基础上增强了: +- torch distributed 初始化与 rank/world_size 协同; +- all_gather 同步控制(例如不同 rank 的 LLM 样本是否齐备); +- 多卡下 WM/LLM 阶段一致性推进与 barrier 同步。 -With proper tuning, PriorZero should achieve: +适用场景:正式大规模实验(推荐使用该入口)。 -- **Exploration Efficiency**: Fewer invalid actions searched (thanks to LLM priors) -- **Sample Efficiency**: Faster convergence (thanks to world model planning) -- **Generalization**: Better performance on unseen games (thanks to LLM knowledge) -- **Memory Efficiency**: No OOM on single GPU (thanks to gradient accumulation) +--- -## 🔍 Monitoring Training +## 2. 主实验启动命令(重点:改哪两个文件) -### TensorBoard +主实验建议通过 DDP 脚本启动: ```bash -tensorboard --logdir=./data_priorzero/ --port=6006 -``` - -**Key metrics to watch**: -- `train/wm_total_loss`: World model total loss -- `train/llm_sft_loss`: LLM supervised fine-tuning loss -- `train/llm_rft_loss`: LLM reinforcement fine-tuning loss -- `train/total_loss`: Combined loss -- `train/wm_grad_norm`: World model gradient norm -- `train/llm_grad_norm`: LLM gradient norm -- `collector_iter/reward_mean`: Average episode reward -- `collector_iter/visit_entropy_mean`: MCTS exploration entropy -- `evaluator_step/reward_mean`: Evaluation reward - -### File Logs - -Check `./data_priorzero/{exp_name}/log/` for: -- Training logs with detailed statistics -- LLM prior statistics (success rate, latency, retry count) -- Game segment statistics (MCTS policies, raw obs, search values) - -### Debug Logs - -During training, you'll see: -``` -[LLM Training] Processing X game segments -[LLM Training] First segment stats: mcts_policies=Y, raw_obs=Z/Z, actions=W -[SEGMENT_DEBUG] raw_obs_text = North of House... -``` - -## 🐛 Troubleshooting - -### OOM (Out of Memory) - -**1. Reduce LLM micro-batch size** (most effective): -```python -llm_micro_batch_size=2 # or even 1 -llm_gradient_accumulation_steps=8 # keep this to maintain effective batch size -``` - -**2. Reduce vLLM memory**: -```python -gpu_memory_utilization=0.2 # Default: 0.3 -``` - -**3. Enable LoRA for LLM**: -```python -use_lora=True -lora_r=8 -lora_alpha=16 -``` - -**4. Reduce world model batch size**: -```python -batch_size=16 # Default: 32 (quick test) -``` - -**5. Reduce prompt length**: -```python -prompt_max_len=512 # Default: 1024 (quick test) -generate_max_len=64 # Default: 128 (quick test) -``` - -**6. Reduce MCTS simulations**: -```python -num_simulations=10 # Default: 25 +cd zoo/jericho/priorzero +bash scripts/run_priorzero_ddp.sh ``` -### LLM Generation Issues +实际跑实验前,主要改两个地方: -**Timeout errors**: -```python -# In priorzero_collector.py -await self._async_get_llm_prior(..., timeout=60.0) # Default: 30.0 -``` - -**vLLM initialization errors**: -- Check CUDA version compatibility -- Ensure `VLLM_USE_V1=1` environment variable (set in entry.py) -- Try reducing `gpu_memory_utilization` - -**Empty raw_obs_text**: -- Fixed! Now properly extracts from `obs['raw_obs_text']` -- Check logs for `[SEGMENT_DEBUG] raw_obs_text = ...` - -### Gradient Errors - -**"element 0 of tensors does not require grad"**: -- Fixed! RFT now properly tracks gradients -- Removed `torch.no_grad()` from RFT forward pass - -### Slow Training - -**1. Use Quick Test Config**: -```python -get_priorzero_config_for_quick_test() # Reduces all resources -``` - -**2. Reduce collector environments**: -```python -collector_env_num=2 # Default: 4 -``` - -**3. Reduce update frequency**: -```python -update_per_collect=5 # Default: 10 -``` - -**4. Reduce game segment length**: -```python -game_segment_length=50 # Default: 200 -``` - -### Buffer/Sampling Issues - -**Double sampling fixed**: -- PriorZeroGameBuffer now caches game_segments -- ~50% faster sampling with no memory overhead +1) `src/priorzero_config.py` +- 修改训练/融合/损失等核心配置(例如 `train_schedule`、`mcts_root_logits_dict`、`advantage_type` 等)。 +- 修改模型预设(`MODEL_CONFIGS`)或 `get_priorzero_config` 中与环境相关的设置。 -## 🔄 Recent Fixes & Improvements +2) `scripts/run_priorzero_ddp.sh` +- 修改 `CUDA_DEVICES`、`NPROC_PER_NODE`、`MASTER_PORT`。 +- 修改 `ENV_ID`、`LLM_MODEL`、`USE_COT`。 +- 确认日志目录 `LOG_DIR`。 -### v2.0.4 (Latest) +建议流程: +- 先在 `priorzero_config.py` 固化实验配置模板; +- 再在 `run_priorzero_ddp.sh` 做“本次任务级”覆写(环境名、卡数、端口等); +- 用日志文件名区分实验版本,便于后续对比。 -✅ **Fixed RFT gradient computation error** -- Removed `torch.no_grad()` from RFT forward pass -- Gradients now properly flow through REINFORCE loss - -✅ **Optimized memory efficiency** -- Implemented micro-batching with gradient accumulation for SFT/RFT -- LLM training processes small chunks (2-4 samples) instead of full batch -- Automatic memory cleanup after each micro-batch -- World model still trains with full batches (no slowdown) - -✅ **Fixed raw_obs_text propagation** -- Enhanced `extract_raw_obs_text()` to prioritize `raw_obs_text` field -- Properly passes raw text from collector to GameSegment -- Now captures actual text: "North of House", "Behind House", etc. - -✅ **Optimized game buffer** -- Eliminated double sampling in `_sample_orig_data()` -- Caches game_segments during sampling (~50% faster) -- Returns game_segments as 3rd element in train_data - -## 📚 References - -### Theoretical Foundations - -1. **AlphaGo/AlphaZero**: Policy-guided MCTS -2. **MuZero**: Model-based RL with learned dynamics -3. **UniZero**: Unified world model for various domains -4. **ORZ (OpenAI)**: LLM fine-tuning for reasoning -5. **REINFORCE**: Policy gradient methods for RL - -### Related Papers - -- **UniZero**: "Unifying World Models via Transformers" -- **MuZero**: "Mastering Atari, Go, Chess and Shogi by Planning with a Learned Model" -- **vLLM**: "Efficient Memory Management for Large Language Model Serving" -- **LoRA**: "Low-Rank Adaptation of Large Language Models" - -## 🤝 Contributing - -This is a research codebase. Contributions are welcome! Key areas for improvement: - -1. **Better LLM prompts**: Improve action ranking quality with CoT reasoning -2. **Reward shaping**: Better credit assignment for RFT -3. **Multi-task learning**: Train on multiple games simultaneously -4. **Efficient MCTS**: Reduce simulation budget via better priors -5. **Dynamic action spaces**: Handle variable action sets across games - -## 📝 Citation - -If you use this code in your research, please cite: +--- -```bibtex -@misc{priorzero2025, - title={PriorZero: LLM-Guided World Model Planning}, - author={PriorZero Team}, - year={2025}, - howpublished={\url{https://github.com/opendilab/LightZero}} -} -``` +## 3. 目前实验结果(阶段性结论) -## 📄 License +当前结果可总结为: +- 在 `detective / zork1 / acorncourt / omniquest` 四个环境中,**LLM/WM 交替训练模式**下,实验均出现比 Unizero更早收敛甚至性能更好的趋势; +- 但在 **LLM 冻结**(或后期 LLM 更新不足)设置下,仍需重点讨论: + - 如何让 PriorZero 在训练后期继续稳定收敛; -This project follows the same license as LightZero (Apache 2.0). --- -**Happy Training! 🚀** - -For questions or issues: -- Open an issue on GitHub: https://github.com/opendilab/LightZero/issues -- Check troubleshooting guide above -- Review log files in `./data_priorzero/{exp_name}/log/` +## 4. 后续可能的改进方向 + +### 4.1 融合方式:先验注入位置再设计 +当前重点在根节点融合(root prior)。可探索: +- 在 MCTS 的**模拟扩展阶段**(非根节点)也注入 LLM 先验; +- 设计“深度相关衰减”策略:树越深,先验权重逐步衰减; +- 对比“仅根节点融合” vs “全树局部融合”的收益与开销。 + +### 4.2 LLM 训练优势函数(advantage)更精细 +可探索更细粒度的 advantage 设计: +- 分阶段 advantage(前期探索导向、后期收敛导向); +- token 级 / action 级加权 advantage; +- 结合 trajectory 置信度、模型不确定度做 adaptive reweight。 + +### 4.3 后期收敛稳定性 +围绕“LLM 冻结后如何继续提升”可尝试: +- 周期性解冻 LLM 的轻量层(如 LoRA 层); +- 在后期降低探索温度、提高价值约束; +- 针对高价值轨迹做重采样,提升有效监督密度。 + +### 4.4 跨域泛化验证:扩展到 Vision 环境 +建议在 Atari 等视觉决策环境上验证 PriorZero 的可迁移性: +- 将当前文本交互任务中的先验融合思路迁移到视觉观测 + 离散动作场景; +- 对比文本环境与视觉环境下,root prior / 扩展阶段先验注入的收益差异; +- 评估在高维观测下,LLM(或多模态模型)与 WM 交替训练的稳定性与样本效率。 + +--- \ No newline at end of file diff --git a/zoo/jericho/priorzero/ablation/llm_as_policy/run_llm_as_policy.py b/zoo/jericho/priorzero/ablation/llm_as_policy/run_llm_as_policy.py new file mode 100644 index 000000000..b5d6c9c5e --- /dev/null +++ b/zoo/jericho/priorzero/ablation/llm_as_policy/run_llm_as_policy.py @@ -0,0 +1,332 @@ +#!/usr/bin/env python3 +"""LLM-as-policy ablation for PriorZero Jericho experiments. + +This baseline keeps the PriorZero prompt/action-prior setup, but removes the +world model, MCTS, replay buffer, and all training. At each environment step it +scores the current valid actions with the frozen LLM and executes the best one. +""" + +from __future__ import annotations + +import argparse +import json +import math +import os +import random +import sys +import time +import contextlib +from collections import deque +from pathlib import Path +from types import SimpleNamespace +from typing import Any, Dict, List, Tuple + +import numpy as np +import torch +import vllm + + +REPO_ROOT = Path(__file__).resolve().parents[2] +LIGHTZERO_ROOT = REPO_ROOT.parents[2] +for path in (REPO_ROOT / "src", LIGHTZERO_ROOT): + path_str = str(path) + if path_str not in sys.path: + sys.path.insert(0, path_str) + +from priorzero_config import get_model_config, get_priorzero_config # noqa: E402 +from priorzero_datafactory import DataProcessor # noqa: E402 +from zoo.jericho.envs.jericho_env import JerichoEnv # noqa: E402 + + +class LocalVLLMActor: + """Minimal adapter matching the DataProcessor vLLM interface.""" + + def __init__(self, model_path: str, tensor_parallel_size: int, max_model_len: int, gpu_memory_utilization: float): + self.requests = [] + self.sampling_params = None + self.llm = vllm.LLM( + model=model_path, + tensor_parallel_size=tensor_parallel_size, + max_model_len=max_model_len, + dtype="bfloat16", + gpu_memory_utilization=gpu_memory_utilization, + trust_remote_code=True, + ) + + def add_requests(self, sampling_params, prompt_token_ids): + from vllm.inputs import TokensPrompt + + self.sampling_params = sampling_params + self.requests = [TokensPrompt(prompt_token_ids=r) for r in prompt_token_ids] + + def get_responses(self): + outputs = self.llm.generate( + prompts=self.requests, + sampling_params=self.sampling_params, + use_tqdm=False, + ) + self.requests = [] + return outputs + + def close(self) -> None: + engine = getattr(self.llm, "llm_engine", None) + engine_core = getattr(engine, "engine_core", None) + if engine_core is not None and hasattr(engine_core, "shutdown"): + engine_core.shutdown() + + +def set_seed(seed: int) -> None: + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(seed) + + +def build_data_processor(llm_cfg, exp_name: str) -> DataProcessor: + vllm_engine = LocalVLLMActor( + model_path=llm_cfg.model_name_or_path, + tensor_parallel_size=llm_cfg.vllm_tensor_parallel_size, + max_model_len=llm_cfg.prompt_max_len + llm_cfg.generate_max_len, + gpu_memory_utilization=llm_cfg.gpu_memory_utilization, + ) + strategy = SimpleNamespace(args=llm_cfg) + return DataProcessor( + rank=0, + world_size=1, + vllm_engine=vllm_engine, + strategy=strategy, + model_path=llm_cfg.model_name_or_path, + exp_name=exp_name, + instance_name="llm_as_policy", + ) + + +def normalize_logprobs(logprobs: Dict[str, float], temperature: float) -> Dict[str, float]: + if not logprobs: + return {} + if temperature <= 1e-8: + best = max(logprobs, key=logprobs.get) + return {k: 0.0 if k == best else float("-inf") for k in logprobs} + + scaled = {k: v / temperature for k, v in logprobs.items()} + max_val = max(scaled.values()) + log_z = math.log(sum(math.exp(v - max_val) for v in scaled.values())) + max_val + return {k: v - log_z for k, v in scaled.items()} + + +def choose_action( + llm_prior: Dict[str, float], + valid_actions: List[str], + temperature: float, + sample: bool, +) -> Tuple[int, str, Dict[str, float]]: + if len(valid_actions) == 0: + return 0, "go", {"go": 1.0} + + filtered = {a: llm_prior[a] for a in valid_actions if a in llm_prior} + if not filtered: + return 0, valid_actions[0], {valid_actions[0]: 1.0} + + norm_logprobs = normalize_logprobs(filtered, temperature) + policy = {a: math.exp(lp) for a, lp in norm_logprobs.items()} + z = sum(policy.values()) + policy = {a: p / z for a, p in policy.items()} if z > 0 else {valid_actions[0]: 1.0} + + if sample: + action_names = list(policy.keys()) + probs = np.array([policy[a] for a in action_names], dtype=np.float64) + probs = probs / probs.sum() + action_name = str(np.random.choice(action_names, p=probs)) + else: + action_name = max(policy, key=policy.get) + + return valid_actions.index(action_name), action_name, policy + + +def run_episode( + seed: int, + env_cfg: Dict[str, Any], + data_processor: DataProcessor, + history_len: int, + temperature: float, + sample: bool, +) -> Dict[str, Any]: + set_seed(seed) + env = JerichoEnv(env_cfg) + env.seed(seed, dynamic_seed=False) + obs = env.reset() + history = deque(maxlen=history_len) + trajectory = [] + total_reward = 0.0 + start = time.time() + + try: + done = False + step = 0 + while not done: + valid_actions = list(obs.get("valid_actions", [])) + history_snapshot = list(history) + prompt = data_processor.get_user_prompt( + history=history_snapshot, + current_obs=obs["raw_obs_text"], + valid_actions=valid_actions, + ) + llm_prior_per_seq, _, _ = data_processor.get_llm_prior( + states=[obs["raw_obs_text"]], + valid_actions_list=[valid_actions], + histories=[history_snapshot], + return_cot=True, + ) + llm_prior = dict(llm_prior_per_seq[0]) + action_idx, action_str, policy = choose_action( + llm_prior=llm_prior, + valid_actions=valid_actions, + temperature=temperature, + sample=sample, + ) + + timestep = env.step(action_idx) + reward = float(timestep.reward) + done = bool(timestep.done) + info = dict(timestep.info) + total_reward += reward + + top_actions = sorted(policy.items(), key=lambda x: x[1], reverse=True)[:5] + trajectory.append( + { + "step": step, + "observation": obs["raw_obs_text"], + "prompt": prompt, + "history_len_cfg": history_len, + "history_len_used": len(history_snapshot), + "prompt_includes_valid_actions": bool( + getattr(data_processor.args.user_prompt_dict, "observation_with_valid_actions", False) + ), + "valid_actions": valid_actions, + "action": info.get("action_str", action_str), + "reward": reward, + "score": float(info.get("score", total_reward)), + "top_policy": [{"action": a, "prob": float(p)} for a, p in top_actions], + "done": done, + } + ) + history.append((obs["raw_obs_text"], info.get("action_str", action_str), reward)) + obs = timestep.obs + step += 1 + + final_score = float(trajectory[-1]["score"]) if trajectory else total_reward + return { + "seed": seed, + "score": final_score, + "total_reward": total_reward, + "steps": len(trajectory), + "duration_sec": time.time() - start, + "trajectory": trajectory, + } + finally: + worker = getattr(env, "_valid_actions_worker", None) + if worker is not None: + with contextlib.suppress(Exception): + worker.close() + env.close() + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="PriorZero ablation: frozen LLM as policy") + parser.add_argument("--env_id", type=str, default="detective.z5") + parser.add_argument("--model", type=str, default="qwen2.5-3b") + parser.add_argument("--seeds", type=int, nargs="+", default=[0, 1]) + parser.add_argument("--history_len", "--his_len", type=int, default=25) + parser.add_argument("--temperature", type=float, default=None) + parser.add_argument("--sample", action="store_true", help="Sample from the LLM action prior instead of greedy argmax.") + parser.add_argument("--use_cot", action="store_true") + parser.add_argument("--output_dir", type=str, default="ablation/llm_as_policy/results") + parser.add_argument("--exp_name", type=str, default=None) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") + + env_name = args.env_id.replace(".z5", "") + exp_name = args.exp_name or f"data_ablation/llm_as_policy/{env_name}_{args.model}_his{args.history_len}" + main_cfg, _, llm_cfg = get_priorzero_config( + env_id=args.env_id, + seed=args.seeds[0], + exp_name=exp_name, + use_cot=args.use_cot, + model_key=args.model, + multi_gpu=False, + ) + model_cfg = get_model_config(args.model) + llm_cfg.enable_rft = False + llm_cfg.enable_world_model = False + llm_cfg.history_length = args.history_len + llm_cfg.vllm_enable_sleep = False + llm_cfg.gpu_memory_utilization = model_cfg["gpu_memory_utilization"] + if args.temperature is not None: + llm_cfg.llm_prior_temperature = args.temperature + + output_dir = Path(args.output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + Path(exp_name, "log").mkdir(parents=True, exist_ok=True) + + data_processor = build_data_processor(llm_cfg=llm_cfg, exp_name=exp_name) + try: + env_cfg = dict(main_cfg.env) + results = [] + + for seed in args.seeds: + result = run_episode( + seed=seed, + env_cfg=env_cfg, + data_processor=data_processor, + history_len=args.history_len, + temperature=llm_cfg.llm_prior_temperature, + sample=args.sample, + ) + results.append(result) + print( + f"[LLM-as-policy] seed={seed} score={result['score']} " + f"steps={result['steps']} duration={result['duration_sec']:.1f}s" + ) + + scores = [r["score"] for r in results] + summary = { + "ablation": "llm_as_policy", + "env_id": args.env_id, + "model": args.model, + "model_path": llm_cfg.model_name_or_path, + "history_len": args.history_len, + "temperature": llm_cfg.llm_prior_temperature, + "sample": args.sample, + "seeds": args.seeds, + "score_mean": float(np.mean(scores)) if scores else 0.0, + "score_std": float(np.std(scores)) if scores else 0.0, + "score_min": float(np.min(scores)) if scores else 0.0, + "score_max": float(np.max(scores)) if scores else 0.0, + "results": results, + } + + timestamp = time.strftime("%Y%m%d_%H%M%S") + output_path = output_dir / f"{args.env_id}_{args.model}_his{args.history_len}_{timestamp}.json" + with output_path.open("w", encoding="utf-8") as f: + json.dump(summary, f, indent=2, ensure_ascii=False) + + print( + "[LLM-as-policy] " + f"mean={summary['score_mean']:.3f} std={summary['score_std']:.3f} " + f"min={summary['score_min']:.3f} max={summary['score_max']:.3f}" + ) + print(f"[LLM-as-policy] results saved to {output_path}") + finally: + vllm_engine = getattr(data_processor, "vllm_engine", None) + if vllm_engine is not None and hasattr(vllm_engine, "close"): + with contextlib.suppress(Exception): + vllm_engine.close() + + +if __name__ == "__main__": + main() diff --git a/zoo/jericho/priorzero/ablation/llm_as_policy/run_llm_as_policy.sh b/zoo/jericho/priorzero/ablation/llm_as_policy/run_llm_as_policy.sh new file mode 100755 index 000000000..3e496075b --- /dev/null +++ b/zoo/jericho/priorzero/ablation/llm_as_policy/run_llm_as_policy.sh @@ -0,0 +1,31 @@ +#!/bin/bash +set -x +set -o pipefail + +PRIORZERO_DIR="/mnt/afs/niuyazhe/workspace/xiongjyu/LightZero/zoo/jericho/priorzero" +PYTHON_BIN="/mnt/afs/niuyazhe/workspace/xiongjyu/envs/rft/bin/python" + +CUDA_DEVICES="${CUDA_DEVICES:-0}" +ENV_ID="${ENV_ID:-detective.z5}" +LLM_MODEL="${LLM_MODEL:-qwen2.5-3b}" +HIS_LEN="${HIS_LEN:-25}" +SEEDS="${SEEDS:-0 1}" +LOG_DIR="${LOG_DIR:-${PRIORZERO_DIR}/data_ablation/run_logs}" + +mkdir -p "${LOG_DIR}" +CURRENT_TIME=$(date +"%Y%m%d_%H%M%S") +LOG_FILE="${LOG_DIR}/llm_as_policy_${ENV_ID}_${LLM_MODEL}_his${HIS_LEN}_${CURRENT_TIME}.txt" + +export CUDA_VISIBLE_DEVICES="${CUDA_DEVICES}" +export PYTHONFAULTHANDLER=1 +export TOKENIZERS_PARALLELISM=false + +cd "${PRIORZERO_DIR}" + +"${PYTHON_BIN}" \ + "${PRIORZERO_DIR}/ablation/llm_as_policy/run_llm_as_policy.py" \ + --env_id "${ENV_ID}" \ + --model "${LLM_MODEL}" \ + --history_len "${HIS_LEN}" \ + --seeds ${SEEDS} \ + 2>&1 | tee "${LOG_FILE}" diff --git a/zoo/jericho/priorzero/ablation/rlft/local_ppo.py b/zoo/jericho/priorzero/ablation/rlft/local_ppo.py new file mode 100644 index 000000000..24bd2add8 --- /dev/null +++ b/zoo/jericho/priorzero/ablation/rlft/local_ppo.py @@ -0,0 +1,453 @@ +from __future__ import annotations + +import gc +import math +import os +from collections import defaultdict +from typing import Dict, Optional, Tuple + +import numpy as np +import torch +import torch.nn as nn +from peft import LoraConfig, TaskType, get_peft_model +from tqdm import tqdm +from transformers import AutoModelForCausalLM, AutoTokenizer +from transformers.integrations.deepspeed import HfDeepSpeedConfig +from transformers.trainer import get_scheduler + + +def masked_mean(tensor: torch.Tensor, mask: Optional[torch.Tensor], dim=None) -> torch.Tensor: + if mask is None: + return tensor.mean(dim=dim) + mask = mask.to(dtype=tensor.dtype, device=tensor.device) + return (tensor * mask).sum(dim=dim) / mask.sum(dim=dim).clamp(min=1.0) + + +def log_probs_from_logits(logits: torch.Tensor, labels: torch.Tensor, temperature: float = 1.0) -> torch.Tensor: + log_probs = torch.log_softmax(logits / temperature, dim=-1) + return log_probs.gather(dim=-1, index=labels.unsqueeze(-1)).squeeze(-1) + + +def entropy_from_logits(logits: torch.Tensor) -> torch.Tensor: + probs = torch.softmax(logits, dim=-1) + log_probs = torch.log_softmax(logits, dim=-1) + return -(probs * log_probs).sum(dim=-1) + + +class FixedKLController: + def __init__(self, kl_coef: float): + self.value = float(kl_coef) + + def update(self, current, n_steps): + return None + + +class RLFTActor(nn.Module): + def __init__( + self, + pretrain: str, + attn_implementation: str, + bf16: bool, + ds_config: Optional[dict], + temperature: float, + train_mode_cfg=None, + enable_value_head: bool = True, + ) -> None: + super().__init__() + if ds_config is not None and ds_config["zero_optimization"]["stage"] == 3: + _ = HfDeepSpeedConfig(ds_config) + + self.temperature = temperature + self.train_mode_cfg = train_mode_cfg if train_mode_cfg is not None else {"mode": "full"} + self.train_mode = self.train_mode_cfg.get("mode", "full") + self.enable_value_head = enable_value_head + self.model = AutoModelForCausalLM.from_pretrained( + pretrain, + trust_remote_code=True, + attn_implementation=attn_implementation, + torch_dtype=torch.bfloat16 if bf16 else "auto", + ) + self.model.config.use_cache = False + + if self.train_mode == "lora": + self.model.enable_input_require_grads() + target_modules = self.train_mode_cfg.get("lora_target_modules") + target_modules = list(target_modules) if target_modules else None + lora_config = LoraConfig( + task_type=TaskType.CAUSAL_LM, + inference_mode=False, + r=self.train_mode_cfg.get("lora_r", 16), + lora_alpha=self.train_mode_cfg.get("lora_alpha", 32), + lora_dropout=self.train_mode_cfg.get("lora_dropout", 0.05), + bias=self.train_mode_cfg.get("lora_bias", "none"), + target_modules=target_modules, + ) + self.model = get_peft_model(self.model, lora_config) + elif self.train_mode != "full": + raise ValueError(f"Unsupported train_mode: {self.train_mode}") + + if enable_value_head: + hidden_size = getattr(self.model.config, "hidden_size", None) + if hidden_size is None: + raise ValueError("Cannot enable value head because model.config.hidden_size is missing.") + self.v_head = nn.Linear(hidden_size, 1) + + def forward( + self, + sequences: torch.LongTensor, + action_mask: torch.Tensor, + attention_mask: torch.Tensor, + return_output: bool = False, + return_entropy: bool = False, + return_values: bool = False, + ): + if return_values and not self.enable_value_head: + raise RuntimeError("return_values=True requires enable_value_head=True.") + + rolled_sequences = torch.roll(sequences, shifts=-1, dims=1) + position_ids = attention_mask.long().cumsum(-1) - 1 + position_ids.masked_fill_(attention_mask == 0, 1) + output = self.model( + sequences, + attention_mask=attention_mask, + position_ids=position_ids, + output_hidden_states=return_values, + ) + logits = output.logits.to(torch.float32) + + if return_entropy: + setattr(output, "entropy", entropy_from_logits(logits)[:, :-1]) + + log_probs = log_probs_from_logits(logits, rolled_sequences, temperature=self.temperature)[:, :-1] + action_log_probs = log_probs[:, -action_mask.shape[1] :] * action_mask.float() + + if return_values: + values = self.v_head(output.hidden_states[-1]).squeeze(-1).to(torch.float32)[:, :-1] + prompt_end_idx = (attention_mask.sum(dim=-1) - action_mask.sum(dim=-1) - 1).clamp(min=0).long() + state_values = values.gather(dim=1, index=prompt_end_idx.unsqueeze(-1)).squeeze(-1) + setattr(output, "state_values", state_values) + + return (action_log_probs, output) if return_output else action_log_probs + + def gradient_checkpointing_enable(self, gradient_checkpointing_kwargs=None): + self.model.gradient_checkpointing_enable(gradient_checkpointing_kwargs=gradient_checkpointing_kwargs) + + +class RLFTReferenceModel: + def __init__(self, strategy, pretrain: str): + self.strategy = strategy + model = RLFTActor( + pretrain=pretrain, + attn_implementation=strategy.args.attn_implementation, + bf16=strategy.args.bf16, + ds_config=strategy.get_ds_eval_config(offload=False), + temperature=strategy.args.temperature, + train_mode_cfg=strategy.args.train_mode_dict, + enable_value_head=False, + ) + self.model = strategy.prepare(model, is_rlhf=True) + self.model.eval() + self.micro_train_batch_size = strategy.args.micro_train_batch_size + + @torch.no_grad() + def forward(self, sequences: torch.Tensor, action_mask: torch.Tensor, attention_mask: torch.Tensor) -> torch.Tensor: + device = torch.cuda.current_device() + outs = [] + chunk_size = max(1, self.micro_train_batch_size) + sequences = sequences.to(device) + attention_mask = attention_mask.to(device) + action_mask = action_mask.to(device) + for i in range(0, sequences.size(0), chunk_size): + outs.append( + self.model( + sequences[i : i + chunk_size], + action_mask=action_mask[i : i + chunk_size], + attention_mask=attention_mask[i : i + chunk_size], + ) + ) + return torch.cat(outs, dim=0) + + +class RLFTPPOTrainer: + def __init__(self, strategy, actor, actor_optim, actor_scheduler, micro_train_batch_size: int): + self.strategy = strategy + self.args = strategy.args + self.actor = actor + self.actor_optim = actor_optim + self.actor_scheduler = actor_scheduler + self.micro_train_batch_size = micro_train_batch_size + self.train_iter = 0 + + def train_batch(self, batch_data: Dict[str, torch.Tensor], kl_ctl: FixedKLController): + device = torch.cuda.current_device() + for k, v in batch_data.items(): + if torch.is_tensor(v): + batch_data[k] = v.to(device) + + all_samples_size = batch_data["input_ids"].size(0) + status_list = [] + metrics_buffer = defaultdict(list) + pbar = tqdm( + range(0, all_samples_size, self.micro_train_batch_size), + desc="RLFT PPO batch", + disable=not self.strategy.is_rank_0(), + ) + acc_grad_steps = self.strategy.accumulated_gradient + + for micro_step, start_idx in enumerate(pbar): + end_idx = min(start_idx + self.micro_train_batch_size, all_samples_size) + micro_batch = { + k: (v[start_idx:end_idx] if torch.is_tensor(v) else v) + for k, v in batch_data.items() + } + micro_batch["log_status"] = batch_data["log_status"][start_idx:end_idx] + + action_log_probs, output = self.actor( + micro_batch["input_ids"], + micro_batch["action_mask"], + attention_mask=micro_batch["attention_mask"], + return_output=True, + return_entropy=True, + return_values=True, + ) + current_action_logprobs = masked_mean( + action_log_probs, + micro_batch["action_mask"], + dim=1, + ) + + actor_loss, clipfrac, clip_ratio, approx_kl = self._policy_loss( + log_probs=current_action_logprobs, + old_log_probs=micro_batch["old_action_log_probs"], + advantages=micro_batch["advantages"], + ) + + if self.args.rft_kl_coef > 0 and micro_batch["ref_action_log_probs"] is not None: + ref_action_logprobs = masked_mean( + micro_batch["ref_action_log_probs"], + micro_batch["action_mask"], + dim=1, + ) + kl_loss = (current_action_logprobs - ref_action_logprobs).mean() + else: + kl_loss = torch.tensor(0.0, device=device) + + value_loss, value_clipfrac = self._value_loss( + values=output.state_values, + returns=micro_batch["returns"], + old_values=micro_batch["old_values"], + ) + entropy = masked_mean(output.entropy[:, -micro_batch["action_mask"].shape[1] :], micro_batch["action_mask"]) + + loss = actor_loss + float(kl_ctl.value) * kl_loss + float(self.args.value_loss_coef) * value_loss + if getattr(self.args, "entropy_loss_coef", 0.0) != 0: + loss -= entropy * self.args.entropy_loss_coef + + self.strategy.backward(loss, self.actor, self.actor_optim) + self.strategy.optimizer_step(self.actor_optim, self.actor, self.actor_scheduler, name="rlft_actor") + + metrics_buffer["policy_loss"].append(actor_loss.detach().float().item()) + metrics_buffer["clipfrac"].append(clipfrac.detach().float().item()) + metrics_buffer["clip_ratio"].append(clip_ratio.detach().float().item()) + metrics_buffer["approx_kl"].append(approx_kl.detach().float().item()) + metrics_buffer["ref_kl"].append(kl_loss.detach().float().item()) + metrics_buffer["value_loss"].append(value_loss.detach().float().item()) + metrics_buffer["value_clipfrac"].append(value_clipfrac.detach().float().item()) + metrics_buffer["entropy"].append(entropy.detach().float().item()) + metrics_buffer["input_length"].append( + (micro_batch["attention_mask"].sum() / micro_batch["attention_mask"].shape[0]).detach().float().item() + - (micro_batch["action_mask"].sum() / micro_batch["action_mask"].shape[0]).detach().float().item() + ) + metrics_buffer["response_length"].append( + (micro_batch["action_mask"].sum() / micro_batch["action_mask"].shape[0]).detach().float().item() + ) + for item in micro_batch["log_status"]: + for k, v in item.items(): + metrics_buffer[k].append(float(v)) + + pbar.set_postfix( + { + "policy_loss": metrics_buffer["policy_loss"][-1], + "value_loss": metrics_buffer["value_loss"][-1], + "iter": self.train_iter, + } + ) + + if ((micro_step + 1) % acc_grad_steps == 0) or ((micro_step + 1) == pbar.total): + self.train_iter += 1 + status = { + "iter": self.train_iter, + "policy_loss": float(np.mean(metrics_buffer["policy_loss"])), + "clipfrac": float(np.mean(metrics_buffer["clipfrac"])), + "clip_ratio": float(np.mean(metrics_buffer["clip_ratio"])), + "approx_kl": float(np.mean(metrics_buffer["approx_kl"])), + "ref_kl": float(np.mean(metrics_buffer["ref_kl"])), + "value_loss": float(np.mean(metrics_buffer["value_loss"])), + "value_clipfrac": float(np.mean(metrics_buffer["value_clipfrac"])), + "entropy": float(np.mean(metrics_buffer["entropy"])), + "input_length_mean": float(np.mean(metrics_buffer["input_length"])), + "response_length_mean": float(np.mean(metrics_buffer["response_length"])), + "valid_action_count_mean": float(np.mean(metrics_buffer["valid_action_count"])), + "value_advantage_mean": float(np.mean(metrics_buffer["value_advantage"])), + "value_advantage_max": float(np.max(metrics_buffer["value_advantage"])), + "value_advantage_min": float(np.min(metrics_buffer["value_advantage"])), + "lr": float(self.actor_scheduler.get_last_lr()[0]), + } + status = self.strategy.all_reduce(status) + status_list.append(status) + metrics_buffer.clear() + + return status_list + + def _policy_loss(self, log_probs, old_log_probs, advantages): + log_ratio = log_probs - old_log_probs + ratio = log_ratio.exp() + surr1 = ratio * advantages + surr2 = ratio.clamp(1 - self.args.eps_clip_low_high[0], 1 + self.args.eps_clip_low_high[1]) * advantages + loss = -torch.min(surr1, surr2).mean() + clipped = ratio.gt(1 + self.args.eps_clip_low_high[1]) | ratio.lt(1 - self.args.eps_clip_low_high[0]) + clipfrac = clipped.float().mean() + clip_ratio = (surr2 < surr1).float().mean() + approx_kl = (-log_ratio.detach()).mean() + return loss, clipfrac, clip_ratio, approx_kl + + def _value_loss(self, values, returns, old_values): + values_clipped = old_values + (values - old_values).clamp( + -float(self.args.value_clip_eps), float(self.args.value_clip_eps) + ) + value_loss_unclipped = (values - returns) ** 2 + value_loss_clipped = (values_clipped - returns) ** 2 + value_loss = 0.5 * torch.max(value_loss_unclipped, value_loss_clipped).mean() + value_clipfrac = (value_loss_clipped > value_loss_unclipped).float().mean() + return value_loss, value_clipfrac + + +class RLFTPolicyModel: + def __init__(self, strategy, pretrain: str, max_steps: Optional[int] = None): + self.strategy = strategy + self.args = strategy.args + self.max_steps = max_steps or int(getattr(self.args, "max_steps", 1_000_000)) + + actor = RLFTActor( + pretrain=pretrain, + attn_implementation=self.args.attn_implementation, + bf16=self.args.bf16, + ds_config=strategy.get_ds_train_config(is_actor=True), + temperature=self.args.temperature, + train_mode_cfg=self.args.train_mode_dict, + enable_value_head=True, + ) + strategy.print(actor) + + self.tokenizer = AutoTokenizer.from_pretrained(pretrain, trust_remote_code=True, padding_side="left") + if self.tokenizer.pad_token is None: + self.tokenizer.pad_token = self.tokenizer.eos_token + + actor_optim = strategy.create_optimizer( + actor, + lr=self.args.learning_rate, + betas=self.args.adam_betas, + weight_decay=self.args.weight_decay, + ) + actor_scheduler = get_scheduler( + self.args.lr_scheduler, + actor_optim, + num_warmup_steps=math.ceil(self.max_steps * self.args.lr_warmup_ratio), + num_training_steps=self.max_steps, + scheduler_specific_kwargs={"min_lr": self.args.learning_rate * 0.1}, + ) + if self.args.gradient_checkpointing: + actor.gradient_checkpointing_enable( + gradient_checkpointing_kwargs={"use_reentrant": self.args.gradient_checkpointing_use_reentrant} + ) + + self.actor, self.actor_optim, self.actor_scheduler = strategy.prepare( + (actor, actor_optim, actor_scheduler), + is_rlhf=True, + ) + self.trainer = RLFTPPOTrainer( + strategy=strategy, + actor=self.actor, + actor_optim=self.actor_optim, + actor_scheduler=self.actor_scheduler, + micro_train_batch_size=self.args.micro_train_batch_size, + ) + self.micro_train_batch_size = self.args.micro_train_batch_size + + def fit(self, batch_data, kl_ctl): + torch.cuda.empty_cache() + self.actor.train() + status = self.trainer.train_batch(batch_data, kl_ctl) + torch.cuda.empty_cache() + torch.cuda.synchronize() + return status + + @torch.no_grad() + def forward_logprobs_values(self, sequences, action_mask, attention_mask) -> Tuple[torch.Tensor, torch.Tensor]: + self.actor.eval() + device = torch.cuda.current_device() + sequences = sequences.to(device) + attention_mask = attention_mask.to(device) + action_mask = action_mask.to(device) + logprob_outs, value_outs = [], [] + chunk_size = max(1, self.micro_train_batch_size) + for i in range(0, sequences.size(0), chunk_size): + log_probs, output = self.actor( + sequences[i : i + chunk_size], + action_mask=action_mask[i : i + chunk_size], + attention_mask=attention_mask[i : i + chunk_size], + return_output=True, + return_values=True, + ) + logprob_outs.append(log_probs) + value_outs.append(output.state_values) + return torch.cat(logprob_outs, dim=0), torch.cat(value_outs, dim=0) + + def save_model(self): + if not self.strategy.is_rank_0(): + return + os.makedirs(self.args.save_path, exist_ok=True) + module = self.actor.module if hasattr(self.actor, "module") else self.actor + module.model.save_pretrained(self.args.save_path) + torch.save(module.v_head.state_dict(), os.path.join(self.args.save_path, "value_head.pt")) + self.tokenizer.save_pretrained(self.args.save_path) + + @property + def train_iter(self): + return self.trainer.train_iter + + +class RLFTTrainer: + def __init__(self, cfg, strategy, policy_model: RLFTPolicyModel, reference_model: Optional[RLFTReferenceModel]): + self.cfg = cfg + self.strategy = strategy + self.policy_model = policy_model + self.reference_model = reference_model + self.kl_ctl = FixedKLController(float(getattr(cfg, "rft_kl_coef", 0.0))) + + def train_batch(self, data, collect_env_steps: int): + ( + input_ids, + attention_mask, + action_mask, + advantages, + rollout_lp, + returns, + old_values, + log_status, + ) = data + batch = { + "input_ids": input_ids, + "attention_mask": attention_mask, + "action_mask": action_mask, + "advantages": advantages, + "old_action_log_probs": rollout_lp, + "returns": returns, + "old_values": old_values, + "log_status": log_status, + } + if self.reference_model is not None: + batch["ref_action_log_probs"] = self.reference_model.forward(input_ids, action_mask, attention_mask) + else: + batch["ref_action_log_probs"] = None + return self.policy_model.fit(batch, self.kl_ctl) diff --git a/zoo/jericho/priorzero/ablation/rlft/run_rlft.py b/zoo/jericho/priorzero/ablation/rlft/run_rlft.py new file mode 100644 index 000000000..ad10d5754 --- /dev/null +++ b/zoo/jericho/priorzero/ablation/rlft/run_rlft.py @@ -0,0 +1,832 @@ +#!/usr/bin/env python3 +"""RLFT ablation for PriorZero Jericho experiments. + +This experiment keeps the PriorZero prompt/action-prior setup, but removes the +world model, MCTS, and replay buffer. It performs the common RLFT loop: + +1. Roll out N full episodes with the current LLM policy in Jericho. +2. Merge all step-level samples into one rollout buffer. +3. Estimate advantages with GAE from a value head on the actor backbone. +4. Run K PPO epochs by sampling minibatches from that rollout buffer. +""" + +from __future__ import annotations + +import argparse +import contextlib +import json +import math +import os +import random +import sys +import time +import gc +from collections import deque +from pathlib import Path +from typing import Any, Dict, List, Tuple + +import numpy as np +import torch +import torch.distributed as dist + + +REPO_ROOT = Path(__file__).resolve().parents[2] +LIGHTZERO_ROOT = REPO_ROOT.parents[2] +for path in (REPO_ROOT / "src", LIGHTZERO_ROOT): + path_str = str(path) + if path_str not in sys.path: + sys.path.insert(0, path_str) + +from priorzero_config import get_priorzero_config # noqa: E402 +from strategy.deepspeed import get_strategy, torch_dist_barrier_and_cuda_sync # noqa: E402 +from zoo.jericho.envs.jericho_env import JerichoEnv # noqa: E402 +from local_ppo import RLFTPolicyModel, RLFTReferenceModel, RLFTTrainer # noqa: E402 + + +def set_seed(seed: int) -> None: + random.seed(seed) + np.random.seed(seed) + torch.manual_seed(seed) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(seed) + + +def normalize_logprobs(logprobs: Dict[str, float], temperature: float) -> Dict[str, float]: + if not logprobs: + return {} + if temperature <= 1e-8: + best = max(logprobs, key=logprobs.get) + return {k: 0.0 if k == best else float("-inf") for k in logprobs} + scaled = {k: v / temperature for k, v in logprobs.items()} + max_val = max(scaled.values()) + log_z = math.log(sum(math.exp(v - max_val) for v in scaled.values())) + max_val + return {k: v - log_z for k, v in scaled.items()} + + +def choose_action( + llm_prior: Dict[str, float], + valid_actions: List[str], + temperature: float, + sample: bool, +) -> Tuple[int, str, Dict[str, float]]: + if len(valid_actions) == 0: + return 0, "go", {"go": 1.0} + + filtered = {a: llm_prior[a] for a in valid_actions if a in llm_prior} + if not filtered: + return 0, valid_actions[0], {valid_actions[0]: 1.0} + + norm_logprobs = normalize_logprobs(filtered, temperature) + policy = {a: math.exp(lp) for a, lp in norm_logprobs.items()} + z = sum(policy.values()) + policy = {a: p / z for a, p in policy.items()} if z > 0 else {valid_actions[0]: 1.0} + + if sample: + action_names = list(policy.keys()) + probs = np.array([policy[a] for a in action_names], dtype=np.float64) + probs = probs / probs.sum() + action_name = str(np.random.choice(action_names, p=probs)) + else: + action_name = max(policy, key=policy.get) + + return valid_actions.index(action_name), action_name, policy + + +class PromptOnlyDataProcessor: + """Use PriorZero prompt utilities without vLLM-backed action scoring.""" + + def __init__(self, llm_cfg, model_path: str): + self.args = llm_cfg + from transformers import AutoTokenizer + + self.tokenizer = AutoTokenizer.from_pretrained( + model_path, trust_remote_code=True, padding_side="left" + ) + if self.tokenizer.pad_token is None: + self.tokenizer.pad_token = self.tokenizer.eos_token + self.use_cot = llm_cfg.use_cot + self.prompt_max_len = llm_cfg.prompt_max_len + self.generate_max_len = llm_cfg.generate_max_len + + def get_system_prompt(self) -> str: + parts = [ + "You are an expert player in a text-based adventure game. Your goal is to maximize the score by choosing the optimal next action.", + "Please analyze the game history and current observation to decide the single best next action.", + "OUTPUT FORMAT:", + ] + if self.use_cot: + parts.append( + "You MUST produce exactly TWO parts in the following order:\n" + "1. Reasoning: Analyze the current situation, available actions, constraints, and uncertainties. Do NOT reveal the final choice here.\n" + "2. Action: The final chosen action.\n" + "Strict Format Example:\n" + "Reasoning: \n" + "Action: " + ) + else: + parts.append( + "Output exactly one line starting with 'Action:'.\n" + "Example:\n" + "Action: " + ) + return "\n".join(parts) + + def get_user_prompt( + self, + history: List[Tuple[str, str, float]] | None = None, + current_obs: str | None = None, + valid_actions: List[str] | None = None, + ) -> str: + prompt_parts = [] + user_prompt_dict = self.args.user_prompt_dict + if history: + prompt_parts.append("=== GAME HISTORY ===") + for i, (obs, action, reward) in enumerate(history, start=1): + prompt_parts.append(f"Step {i}:") + prompt_parts.append(f"Observation: {obs.strip()}") + prompt_parts.append(f"Action: {action.strip()}") + if user_prompt_dict.history_with_reward: + prompt_parts.append(f"Reward: {reward}") + prompt_parts.append("") + + prompt_parts.append("=== CURRENT OBSERVATION ===") + prompt_parts.append((current_obs or "").strip()) + if user_prompt_dict.observation_with_valid_actions and valid_actions: + actions_str = ", ".join([f"'{act}'" for act in valid_actions]) + prompt_parts.append(f"\n[Valid Actions]\nYou can choose from the following actions: {actions_str}") + + prompt_parts.append("\n=== INSTRUCTION ===") + if self.use_cot: + prompt_parts.append( + "Please analyze the situation and provide your response in the following format:\n" + "Reasoning: \n" + "Action: " + ) + else: + prompt_parts.append( + "Decide on the best next move and output it in the following format:\n" + "Action: " + ) + return "\n".join(prompt_parts) + + def build_chat_context(self, user_prompt: str) -> str: + return self.tokenizer.apply_chat_template( + [ + {"role": "system", "content": self.get_system_prompt()}, + {"role": "user", "content": user_prompt}, + ], + tokenize=False, + add_generation_prompt=True, + ) + + +@torch.no_grad() +def score_valid_actions_with_actor( + policy_model: RLFTPolicyModel, + tokenizer, + data_processor: PromptOnlyDataProcessor, + prompt: str, + valid_actions: List[str], + prompt_max_len: int, +) -> Dict[str, Dict[str, Any]]: + if not valid_actions: + valid_actions = ["go"] + + all_context_texts = [data_processor.build_chat_context(prompt) for _ in valid_actions] + context_ids = tokenizer( + all_context_texts, + add_special_tokens=False, + max_length=prompt_max_len - 64, + padding=False, + truncation=True, + )["input_ids"] + label_texts = ["Action: " + action + tokenizer.eos_token for action in valid_actions] + label_ids = tokenizer(label_texts, add_special_tokens=False, padding=False, truncation=False)["input_ids"] + full_ids = [c + l for c, l in zip(context_ids, label_ids)] + + inputs = tokenizer.pad({"input_ids": full_ids}, padding=True, return_tensors="pt") + max_tgt_len = max(len(ids) for ids in label_ids) + action_mask = torch.zeros((len(valid_actions), max_tgt_len), dtype=torch.long) + for idx, ids in enumerate(label_ids): + action_mask[idx, -len(ids):] = 1 + + log_probs, state_values = policy_model.forward_logprobs_values( + sequences=inputs.input_ids, + action_mask=action_mask, + attention_mask=inputs.attention_mask, + ) + log_probs_cpu = log_probs.detach().cpu() + state_values_cpu = state_values.detach().cpu() + action_mask_cpu = action_mask.cpu() + + scored = {} + for idx, action in enumerate(valid_actions): + lp_tokens = log_probs_cpu[idx, action_mask_cpu[idx].bool()].tolist() + score = float(sum(lp_tokens) / max(len(lp_tokens), 1)) + scored[action] = { + "score": score, + "rollout_logprob": lp_tokens, + "full_ids": full_ids[idx], + "label_ids": label_ids[idx], + "value": float(state_values_cpu[idx].item()), + } + return scored + + +def close_env(env: JerichoEnv) -> None: + worker = getattr(env, "_valid_actions_worker", None) + if worker is not None: + with contextlib.suppress(Exception): + worker.close() + env.close() + + +def compute_gae( + rewards: List[float], + values: List[float], + dones: List[bool], + gamma: float, + gae_lambda: float, +) -> Tuple[List[float], List[float]]: + advantages = [0.0 for _ in rewards] + last_gae = 0.0 + for t in reversed(range(len(rewards))): + if t == len(rewards) - 1: + next_non_terminal = 0.0 if dones[t] else 1.0 + next_value = 0.0 + else: + next_non_terminal = 0.0 if dones[t] else 1.0 + next_value = values[t + 1] + delta = rewards[t] + gamma * next_value * next_non_terminal - values[t] + last_gae = delta + gamma * gae_lambda * next_non_terminal * last_gae + advantages[t] = last_gae + returns = [adv + value for adv, value in zip(advantages, values)] + return advantages, returns + + +def rollout_episode( + env_cfg: Dict[str, Any], + data_processor: PromptOnlyDataProcessor, + seed: int, + history_len: int, + temperature: float, + sample: bool, + policy_model: RLFTPolicyModel, +) -> Dict[str, Any]: + set_seed(seed) + env = JerichoEnv(env_cfg) + env.seed(seed, dynamic_seed=False) + obs = env.reset() + history = deque(maxlen=history_len) + samples = [] + trajectory = [] + rewards = [] + total_reward = 0.0 + start = time.time() + + try: + done = False + step = 0 + while not done: + valid_actions = list(obs.get("valid_actions", [])) + if not valid_actions: + valid_actions = ["go"] + history_snapshot = list(history) + prompt = data_processor.get_user_prompt( + history=history_snapshot, + current_obs=obs["raw_obs_text"], + valid_actions=valid_actions, + ) + scored_actions = score_valid_actions_with_actor( + policy_model=policy_model, + tokenizer=data_processor.tokenizer, + data_processor=data_processor, + prompt=prompt, + valid_actions=valid_actions, + prompt_max_len=data_processor.prompt_max_len, + ) + llm_prior = {action: info["score"] for action, info in scored_actions.items()} + action_idx, action_str, policy = choose_action( + llm_prior=llm_prior, + valid_actions=valid_actions, + temperature=temperature, + sample=sample, + ) + norm_logprobs = normalize_logprobs(llm_prior, temperature) + + action_info = scored_actions[action_str] + if len(action_info["label_ids"]) > 0: + candidate_actions = [a for a in valid_actions if a in scored_actions] + samples.append( + { + "prompt": prompt, + "action": action_str, + "valid_actions": candidate_actions, + "full_ids": action_info["full_ids"], + "label_ids": action_info["label_ids"], + "candidate_full_ids": [scored_actions[a]["full_ids"] for a in candidate_actions], + "candidate_label_ids": [scored_actions[a]["label_ids"] for a in candidate_actions], + "chosen_candidate_index": candidate_actions.index(action_str), + "old_action_logprob": float(action_info["score"]), + "old_categorical_logprob": float(norm_logprobs[action_str]), + "value": action_info["value"], + } + ) + + timestep = env.step(action_idx) + reward = float(timestep.reward) + done = bool(timestep.done) + info = dict(timestep.info) + if samples: + samples[-1]["reward"] = reward + samples[-1]["done"] = done + rewards.append(reward) + total_reward += reward + + top_actions = sorted(policy.items(), key=lambda x: x[1], reverse=True)[:5] + trajectory.append( + { + "step": step, + "observation": obs["raw_obs_text"], + "prompt": prompt, + "history_len_cfg": history_len, + "history_len_used": len(history_snapshot), + "prompt_includes_valid_actions": bool( + getattr(data_processor.args.user_prompt_dict, "observation_with_valid_actions", False) + ), + "valid_actions": valid_actions, + "action": info.get("action_str", action_str), + "reward": reward, + "score": float(info.get("score", total_reward)), + "top_policy": [{"action": a, "prob": float(p)} for a, p in top_actions], + "done": done, + } + ) + history.append((obs["raw_obs_text"], info.get("action_str", action_str), reward)) + obs = timestep.obs + step += 1 + + final_score = float(trajectory[-1]["score"]) if trajectory else total_reward + return { + "seed": seed, + "score": final_score, + "total_reward": total_reward, + "steps": len(trajectory), + "duration_sec": time.time() - start, + "samples": samples, + "trajectory": trajectory, + } + finally: + close_env(env) + + +def gather_advantage_stats(local_advantages: List[float], world_size: int) -> Tuple[float, float]: + if world_size <= 1: + arr = np.asarray(local_advantages, dtype=np.float32) + else: + gathered = [None for _ in range(world_size)] + dist.all_gather_object(gathered, local_advantages) + arr = np.asarray([x for rank_values in gathered for x in rank_values], dtype=np.float32) + if arr.size == 0: + return 0.0, 1.0 + return float(arr.mean()), float(arr.std() + 1e-8) + + +def attach_gae( + rollouts: List[Dict[str, Any]], + samples: List[Dict[str, Any]], + gamma: float, + gae_lambda: float, +) -> None: + if not samples: + return + + cursor = 0 + for rollout in rollouts: + ep_samples = rollout["samples"] + ep_n = len(ep_samples) + ep_values = [float(s["value"]) for s in ep_samples] + ep_rewards = [float(s["reward"]) for s in ep_samples] + ep_dones = [bool(s["done"]) for s in ep_samples] + ep_advantages, ep_returns = compute_gae( + rewards=ep_rewards, + values=ep_values, + dones=ep_dones, + gamma=gamma, + gae_lambda=gae_lambda, + ) + for sample_item, value, advantage, ret in zip(ep_samples, ep_values, ep_advantages, ep_returns): + sample_item["value"] = value + sample_item["advantage"] = float(advantage) + sample_item["return"] = float(ret) + cursor += ep_n + + +def build_train_batch( + samples: List[Dict[str, Any]], + tokenizer, + adv_mean: float, + adv_std: float, +) -> Tuple[ + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, + torch.Tensor, + List[Dict[str, float]], +]: + if not samples: + raise ValueError("Cannot build RLFT batch from empty samples.") + + full_ids_list = [s["full_ids"] for s in samples] + label_ids_list = [s["label_ids"] for s in samples] + inputs = tokenizer.pad({"input_ids": full_ids_list}, padding=True, return_tensors="pt") + max_tgt_len = max(len(ids) for ids in label_ids_list) + action_mask = torch.zeros((len(samples), max_tgt_len), dtype=torch.long) + for idx, ids in enumerate(label_ids_list): + action_mask[idx, -len(ids):] = 1 + + old_action_logprobs = torch.zeros((len(samples),), dtype=torch.float32) + returns = torch.zeros((len(samples),), dtype=torch.float32) + old_values = torch.zeros((len(samples),), dtype=torch.float32) + + normalized_advantages = [] + log_status = [] + for idx, sample_item in enumerate(samples): + raw_adv = float(sample_item["advantage"]) + norm_adv = (raw_adv - adv_mean) / adv_std + old_action_logprobs[idx] = float(sample_item["old_action_logprob"]) + returns[idx] = float(sample_item["return"]) + old_values[idx] = float(sample_item["value"]) + normalized_advantages.append(norm_adv) + log_status.append( + { + "value_advantage": float(norm_adv), + "raw_gae_advantage": raw_adv, + "value_target_return": float(sample_item["return"]), + "old_value": float(sample_item["value"]), + "valid_action_count": float(len(sample_item.get("valid_actions", []))), + } + ) + + return ( + inputs.input_ids, + inputs.attention_mask, + action_mask, + torch.tensor(normalized_advantages, dtype=torch.float32), + old_action_logprobs, + returns, + old_values, + log_status, + ) + + +def memory_cleanup() -> None: + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + + +def iter_fixed_count_minibatches( + samples: List[Dict[str, Any]], + minibatch_size: int, + num_minibatches: int, + rng: random.Random, +): + epoch_samples = list(samples) + rng.shuffle(epoch_samples) + for batch_idx in range(num_minibatches): + start_idx = batch_idx * minibatch_size + minibatch = epoch_samples[start_idx : start_idx + minibatch_size] + if len(minibatch) < minibatch_size: + minibatch = minibatch + rng.choices(samples, k=minibatch_size - len(minibatch)) + yield minibatch + + +def write_jsonl(path: Path, record: Dict[str, Any]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with path.open("a", encoding="utf-8") as f: + f.write(json.dumps(record, ensure_ascii=False) + "\n") + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="PriorZero ablation: RLFT without world model") + parser.add_argument("--env_id", type=str, default="detective.z5") + parser.add_argument("--model", type=str, default="qwen2.5-3b") + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--history_len", "--his_len", type=int, default=25) + parser.add_argument("--max_env_steps", type=int, default=100_000) + parser.add_argument("--max_rlft_iters", type=int, default=1_000_000) + parser.add_argument("--rollout_episodes_per_iter", type=int, default=50) + parser.add_argument("--ppo_epochs", type=int, default=2) + parser.add_argument("--ppo_minibatch_size", type=int, default=128) + parser.add_argument("--eval_episodes", type=int, default=2) + parser.add_argument("--eval_freq", type=int, default=1) + parser.add_argument("--gamma", type=float, default=1.0) + parser.add_argument("--gae_lambda", type=float, default=0.95) + parser.add_argument("--temperature", type=float, default=None) + parser.add_argument("--micro_train_batch_size", type=int, default=1) + parser.add_argument("--learning_rate", type=float, default=1e-6) + parser.add_argument("--kl_coef", type=float, default=0.01) + parser.add_argument("--value_loss_coef", type=float, default=0.5) + parser.add_argument("--value_clip_eps", type=float, default=0.2) + parser.add_argument("--zero_stage", type=int, default=2) + parser.add_argument("--use_cot", action="store_true") + parser.add_argument("--sample_eval", action="store_true") + parser.add_argument("--output_dir", type=str, default="ablation/rlft/results") + parser.add_argument("--exp_name", type=str, default=None) + parser.add_argument("--save_freq", type=int, default=10) + return parser.parse_args() + + +def main() -> None: + args = parse_args() + if args.rollout_episodes_per_iter <= 0: + raise ValueError("--rollout_episodes_per_iter must be positive.") + if args.ppo_epochs <= 0: + raise ValueError("--ppo_epochs must be positive.") + if args.ppo_minibatch_size <= 0: + raise ValueError("--ppo_minibatch_size must be positive.") + if args.max_env_steps <= 0: + raise ValueError("--max_env_steps must be positive.") + if args.zero_stage != 2: + raise ValueError("This ablation is configured for DeepSpeed ZeRO-2; use --zero_stage 2.") + os.environ.setdefault("TOKENIZERS_PARALLELISM", "false") + + rank = int(os.environ.get("RANK", "0")) + local_rank = int(os.environ.get("LOCAL_RANK", "0")) + if torch.cuda.is_available(): + torch.cuda.set_device(local_rank) + + env_name = args.env_id.replace(".z5", "") + exp_name = args.exp_name or f"data_ablation/rlft/{env_name}_{args.model}_his{args.history_len}" + main_cfg, _, llm_cfg = get_priorzero_config( + env_id=args.env_id, + seed=args.seed, + exp_name=exp_name, + use_cot=args.use_cot, + model_key=args.model, + multi_gpu=True, + ) + llm_cfg.enable_world_model = False + llm_cfg.enable_rft = True + llm_cfg.history_length = args.history_len + llm_cfg.train_batch_size = args.ppo_minibatch_size + llm_cfg.micro_train_batch_size = args.micro_train_batch_size + llm_cfg.learning_rate = args.learning_rate + llm_cfg.zero_stage = args.zero_stage + llm_cfg.rft_kl_coef = args.kl_coef + llm_cfg.enable_value_head = True + llm_cfg.value_loss_coef = args.value_loss_coef + llm_cfg.value_clip_eps = args.value_clip_eps + llm_cfg.rlft_action_temperature = llm_cfg.llm_prior_temperature + llm_cfg.policy_loss_type = "ppo" + llm_cfg.use_rollout_as_old_policy = True + llm_cfg.enable_vllm_is_correction = False + llm_cfg.vllm_enable_sleep = False + llm_cfg.enable_vllm = False + llm_cfg.disable_vllm_sync = True + llm_cfg.max_steps = args.max_rlft_iters + llm_cfg.seed = args.seed + llm_cfg.llm_save_freq = args.save_freq + llm_cfg.save_path = f"./{exp_name}/llm_ckpt/" + if args.temperature is not None: + llm_cfg.llm_prior_temperature = args.temperature + llm_cfg.rlft_action_temperature = args.temperature + + strategy = get_strategy(llm_cfg) + strategy.setup_distributed() + world_size = strategy.world_size + if args.rollout_episodes_per_iter < world_size: + raise ValueError( + f"--rollout_episodes_per_iter ({args.rollout_episodes_per_iter}) must be >= world_size ({world_size})." + ) + if args.ppo_minibatch_size < world_size: + raise ValueError( + f"--ppo_minibatch_size ({args.ppo_minibatch_size}) must be >= world_size ({world_size})." + ) + if args.ppo_minibatch_size < args.micro_train_batch_size * world_size: + raise ValueError( + "--ppo_minibatch_size must be at least micro_train_batch_size * world_size " + f"({args.micro_train_batch_size * world_size})." + ) + local_ppo_minibatch_size = max(1, args.ppo_minibatch_size // world_size) + set_seed(args.seed + rank) + + Path(exp_name, "log").mkdir(parents=True, exist_ok=True) + output_dir = Path(args.output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + rollout_log_path = output_dir / f"{args.env_id}_{args.model}_his{args.history_len}_seed{args.seed}_rollouts.jsonl" + eval_log_path = output_dir / f"{args.env_id}_{args.model}_his{args.history_len}_seed{args.seed}_eval.jsonl" + + vllm_engine = None + data_processor = PromptOnlyDataProcessor(llm_cfg=llm_cfg, model_path=llm_cfg.model_name_or_path) + + ref_model = RLFTReferenceModel(strategy=strategy, pretrain=llm_cfg.model_name_or_path) if llm_cfg.rft_kl_coef > 0 else None + policy_model = RLFTPolicyModel( + strategy=strategy, + pretrain=llm_cfg.model_name_or_path, + max_steps=llm_cfg.max_steps, + ) + trainer = RLFTTrainer( + cfg=llm_cfg, + strategy=strategy, + policy_model=policy_model, + reference_model=ref_model, + ) + + env_cfg = dict(main_cfg.env) + torch_dist_barrier_and_cuda_sync() + total_env_steps = 0 + total_episodes = 0 + + train_iter = 0 + while train_iter < args.max_rlft_iters and total_env_steps < args.max_env_steps: + local_rollouts = [] + local_samples = [] + for episode_idx in range(args.rollout_episodes_per_iter): + if episode_idx % world_size != rank: + continue + rollout_seed = args.seed + train_iter * args.rollout_episodes_per_iter + episode_idx + rollout = rollout_episode( + env_cfg=env_cfg, + data_processor=data_processor, + seed=rollout_seed, + history_len=args.history_len, + temperature=llm_cfg.llm_prior_temperature, + sample=True, + policy_model=policy_model, + ) + local_samples.extend(rollout["samples"]) + local_rollouts.append(rollout) + + attach_gae( + rollouts=local_rollouts, + samples=local_samples, + gamma=args.gamma, + gae_lambda=args.gae_lambda, + ) + local_advantages = [float(s["advantage"]) for s in local_samples] + adv_mean, adv_std = gather_advantage_stats(local_advantages, world_size=world_size) + + local_ready = int(len(local_samples) > 0) + ready_flags = [None for _ in range(world_size)] + dist.all_gather_object(ready_flags, local_ready) + if min(ready_flags) == 0: + if rank == 0: + print(f"[RLFT] skip train_iter={train_iter}, ready_flags={ready_flags}") + train_iter += 1 + continue + + gathered_rollouts = [None for _ in range(world_size)] + dist.all_gather_object(gathered_rollouts, local_rollouts) + if rank == 0: + flat_rollouts = [r for rank_rollouts in gathered_rollouts for r in rank_rollouts] + iter_env_steps = int(sum(r["steps"] for r in flat_rollouts)) + iter_episodes = int(len(flat_rollouts)) + total_env_steps += iter_env_steps + total_episodes += iter_episodes + else: + flat_rollouts = None + iter_env_steps = 0 + iter_episodes = 0 + counters = [total_env_steps, total_episodes, iter_env_steps, iter_episodes] + dist.broadcast_object_list(counters, src=0) + total_env_steps, total_episodes, iter_env_steps, iter_episodes = [int(x) for x in counters] + + local_minibatches = int(math.ceil(len(local_samples) / local_ppo_minibatch_size)) + gathered_minibatches = [None for _ in range(world_size)] + dist.all_gather_object(gathered_minibatches, local_minibatches) + minibatches_per_epoch = int(max(gathered_minibatches)) + + train_rng = random.Random(args.seed + train_iter * 1_000_003 + rank) + train_statuses = [] + minibatches_trained = 0 + optimizer_updates_before = int(getattr(policy_model, "train_iter", 0)) + for ppo_epoch in range(args.ppo_epochs): + for minibatch_samples in iter_fixed_count_minibatches( + samples=local_samples, + minibatch_size=local_ppo_minibatch_size, + num_minibatches=minibatches_per_epoch, + rng=train_rng, + ): + batch = build_train_batch( + samples=minibatch_samples, + tokenizer=data_processor.tokenizer, + adv_mean=adv_mean, + adv_std=adv_std, + ) + status = trainer.train_batch(batch, collect_env_steps=total_env_steps) + train_statuses.extend(status or []) + minibatches_trained += 1 + memory_cleanup() + optimizer_updates_after = int(getattr(policy_model, "train_iter", optimizer_updates_before)) + optimizer_updates = optimizer_updates_after - optimizer_updates_before + + if rank == 0: + scores = [r["score"] for r in flat_rollouts] + record = { + "iter": train_iter, + "phase": "train_rollout", + "env_steps": total_env_steps, + "episode_count": total_episodes, + "iter_env_steps": iter_env_steps, + "iter_episode_count": iter_episodes, + "scores": scores, + "episode_return_mean": float(np.mean(scores)) if scores else 0.0, + "episode_return_std": float(np.std(scores)) if scores else 0.0, + "episode_return_min": float(np.min(scores)) if scores else 0.0, + "episode_return_max": float(np.max(scores)) if scores else 0.0, + "advantage_mean": adv_mean, + "advantage_std": adv_std, + "rollout_sample_count": int(sum(len(r["samples"]) for r in flat_rollouts)), + "local_train_sample_count": len(local_samples), + "ppo_epochs": args.ppo_epochs, + "ppo_minibatch_size": args.ppo_minibatch_size, + "local_ppo_minibatch_size": local_ppo_minibatch_size, + "minibatches_per_epoch": minibatches_per_epoch, + "ppo_minibatches_trained": minibatches_trained, + "optimizer_updates": optimizer_updates, + "last_train_status": train_statuses[-1] if train_statuses else {}, + "gamma": args.gamma, + "gae_lambda": args.gae_lambda, + "max_env_steps": args.max_env_steps, + "history_len": args.history_len, + "prompt_includes_valid_actions": True, + "rollouts": flat_rollouts, + } + write_jsonl(rollout_log_path, record) + print( + f"[RLFT] iter={train_iter} env_steps={total_env_steps} episodes={total_episodes} " + f"iter_episodes={iter_episodes} iter_env_steps={iter_env_steps} " + f"episode_return_mean={record['episode_return_mean']:.3f} " + f"episode_return_min={record['episode_return_min']:.3f} " + f"episode_return_max={record['episode_return_max']:.3f} " + f"samples={record['rollout_sample_count']} ppo_epochs={args.ppo_epochs} " + f"minibatches_per_epoch={minibatches_per_epoch} " + f"ppo_minibatches={minibatches_trained} optimizer_updates={optimizer_updates} " + f"adv_mean={adv_mean:.3f} adv_std={adv_std:.3f}" + ) + + if args.eval_freq > 0 and (train_iter % args.eval_freq == 0 or train_iter == args.max_rlft_iters - 1): + local_eval_rollouts = [] + for ep in range(args.eval_episodes): + if ep % world_size != rank: + continue + eval_seed = args.seed + 10_000 + train_iter * args.eval_episodes + ep + eval_rollout = rollout_episode( + env_cfg=env_cfg, + data_processor=data_processor, + seed=eval_seed, + history_len=args.history_len, + temperature=llm_cfg.llm_prior_temperature, + sample=args.sample_eval, + policy_model=policy_model, + ) + eval_rollout.pop("samples") + local_eval_rollouts.append(eval_rollout) + gathered_eval = [None for _ in range(world_size)] + dist.all_gather_object(gathered_eval, local_eval_rollouts) + if rank == 0: + flat_eval = [r for rank_rollouts in gathered_eval for r in rank_rollouts] + scores = [r["score"] for r in flat_eval] + eval_env_steps = int(sum(r["steps"] for r in flat_eval)) + record = { + "iter": train_iter, + "phase": "eval", + "env_steps": total_env_steps, + "episode_count": total_episodes, + "eval_env_steps": eval_env_steps, + "eval_episode_count": int(len(flat_eval)), + "scores": scores, + "episode_return_mean": float(np.mean(scores)) if scores else 0.0, + "episode_return_std": float(np.std(scores)) if scores else 0.0, + "episode_return_min": float(np.min(scores)) if scores else 0.0, + "episode_return_max": float(np.max(scores)) if scores else 0.0, + "history_len": args.history_len, + "prompt_includes_valid_actions": True, + "rollouts": flat_eval, + } + write_jsonl(eval_log_path, record) + print( + f"[RLFT][eval] iter={train_iter} env_steps={total_env_steps} episodes={total_episodes} " + f"eval_episodes={len(flat_eval)} eval_env_steps={eval_env_steps} " + f"episode_return_mean={record['episode_return_mean']:.3f} scores={scores}" + ) + + torch_dist_barrier_and_cuda_sync() + train_iter += 1 + + policy_model.save_model() + if rank == 0: + print(f"[RLFT] finished. rollout_log={rollout_log_path}, eval_log={eval_log_path}") + + torch_dist_barrier_and_cuda_sync() + if dist.is_initialized(): + dist.destroy_process_group() + + +if __name__ == "__main__": + main() diff --git a/zoo/jericho/priorzero/ablation/rlft/run_rlft.sh b/zoo/jericho/priorzero/ablation/rlft/run_rlft.sh new file mode 100755 index 000000000..9a62a6618 --- /dev/null +++ b/zoo/jericho/priorzero/ablation/rlft/run_rlft.sh @@ -0,0 +1,46 @@ +#!/bin/bash +set -e +set -x +set -o pipefail + +PRIORZERO_DIR="/mnt/afs/niuyazhe/workspace/xiongjyu/LightZero/zoo/jericho/priorzero" + +CUDA_DEVICES="${CUDA_DEVICES:-0}" +NPROC_PER_NODE="${NPROC_PER_NODE:-1}" +MASTER_PORT="${MASTER_PORT:-24564}" + +ENV_ID="${ENV_ID:-detective.z5}" +LLM_MODEL="${LLM_MODEL:-qwen2.5-3b}" +HIS_LEN="${HIS_LEN:-25}" +SEEDS="${SEEDS:-0 1}" +MAX_ENV_STEPS="${MAX_ENV_STEPS:-100000}" +ROLLOUT_EPISODES_PER_ITER="${ROLLOUT_EPISODES_PER_ITER:-50}" +LOG_DIR="${LOG_DIR:-${PRIORZERO_DIR}/data_ablation/run_logs}" + +mkdir -p "${LOG_DIR}" + +export CUDA_VISIBLE_DEVICES="${CUDA_DEVICES}" +export PYTHONFAULTHANDLER=1 +export TOKENIZERS_PARALLELISM=false +export TORCH_DISTRIBUTED_DEBUG="${TORCH_DISTRIBUTED_DEBUG:-OFF}" +export NCCL_DEBUG="${NCCL_DEBUG:-WARN}" + +cd "${PRIORZERO_DIR}" + +for SEED in ${SEEDS}; do + CURRENT_TIME=$(date +"%Y%m%d_%H%M%S") + LOG_FILE="${LOG_DIR}/rlft_${ENV_ID}_${LLM_MODEL}_his${HIS_LEN}_seed${SEED}_${CURRENT_TIME}.txt" + RUN_MASTER_PORT=$((MASTER_PORT + SEED)) + + torchrun \ + --nproc_per_node="${NPROC_PER_NODE}" \ + --master-port="${RUN_MASTER_PORT}" \ + "${PRIORZERO_DIR}/ablation/rlft/run_rlft.py" \ + --env_id "${ENV_ID}" \ + --model "${LLM_MODEL}" \ + --seed "${SEED}" \ + --history_len "${HIS_LEN}" \ + --max_env_steps "${MAX_ENV_STEPS}" \ + --rollout_episodes_per_iter "${ROLLOUT_EPISODES_PER_ITER}" \ + 2>&1 | tee "${LOG_FILE}" +done diff --git a/zoo/jericho/priorzero/async_training_coordinator.py b/zoo/jericho/priorzero/async_training_coordinator.py deleted file mode 100644 index 46a7a36bd..000000000 --- a/zoo/jericho/priorzero/async_training_coordinator.py +++ /dev/null @@ -1,390 +0,0 @@ -# async_training_coordinator.py -""" -[PRIORZERO] Async Training Coordinator - -This module implements async coordination for collect/train/eval tasks. - -Key Features: -- Configurable off-policy degree to control async level -- Automatic fallback to synchronous mode (off_policy_degree=0) -- Independent async evaluation -- Thread-safe buffer access control - -Author: PriorZero Team -Date: 2025-01-21 -""" - -import asyncio -import time -from typing import Optional, Dict, Any, Callable, Awaitable -from loguru import logger - - -class AsyncTrainingCoordinator: - """ - Coordinates async execution of collect, train, and eval tasks. - - The coordinator manages the async execution based on off_policy_degree: - - off_policy_degree = 0: Synchronous mode (collect -> train -> eval) - - off_policy_degree > 0: Async mode with bounded lag - - The off_policy_degree controls how many batches the training can lag - behind the collection. Higher values allow more async execution but - increase off-policy bias. - """ - - def __init__( - self, - off_policy_degree: int = 0, - enable_async_eval: bool = False, - buffer_size: int = 10000, - batch_size: int = 32, - ): - """ - Initialize AsyncTrainingCoordinator. - - Args: - off_policy_degree: Degree of async between collect and train - - 0: Synchronous mode - - >0: Max number of batches train can lag behind collect - - -1: Auto-tune based on buffer_size and batch_size - enable_async_eval: Whether to run eval asynchronously - buffer_size: Replay buffer size (for auto-tuning) - batch_size: Training batch size (for auto-tuning) - """ - self.off_policy_degree = off_policy_degree - self.enable_async_eval = enable_async_eval - self.buffer_size = buffer_size - self.batch_size = batch_size - - # Auto-tune off_policy_degree if set to -1 - if self.off_policy_degree == -1: - # Auto-tune: allow lag up to 10% of buffer capacity - self.off_policy_degree = max(1, (buffer_size // batch_size) // 10) - logger.info(f"Auto-tuned off_policy_degree to {self.off_policy_degree}") - - # Synchronization primitives - self._collect_count = 0 # Number of collect iterations completed - self._train_count = 0 # Number of train iterations completed - self._eval_task: Optional[asyncio.Task] = None - - # Locks for thread-safe access - self._lock = asyncio.Lock() - - # Performance tracking - self._collect_times = [] - self._train_times = [] - self._eval_times = [] - - logger.info(f"AsyncTrainingCoordinator initialized:") - logger.info(f" - off_policy_degree: {self.off_policy_degree}") - logger.info(f" - enable_async_eval: {self.enable_async_eval}") - logger.info(f" - mode: {'SYNCHRONOUS' if self.is_synchronous else 'ASYNCHRONOUS'}") - - @property - def is_synchronous(self) -> bool: - """Check if coordinator is in synchronous mode.""" - return self.off_policy_degree == 0 - - @property - def collect_train_lag(self) -> int: - """Get current lag between collect and train iterations.""" - return self._collect_count - self._train_count - - def can_train(self) -> bool: - """ - Check if training is allowed based on off_policy_degree. - - In synchronous mode (off_policy_degree=0), training must wait for collect. - In async mode, training can proceed as long as lag is within bounds. - """ - if self.is_synchronous: - # Synchronous: train only after collect - return self._collect_count > self._train_count - else: - # Async: train can proceed if there's data and lag is acceptable - # We allow training as long as there's collected data - return self._collect_count > 0 - - def can_collect(self) -> bool: - """ - Check if collection is allowed based on off_policy_degree. - - In synchronous mode, collection must wait for train to finish. - In async mode, collection can proceed as long as lag doesn't exceed limit. - """ - if self.is_synchronous: - # Synchronous: collect only after train - return self._train_count >= self._collect_count - else: - # Async: collect can proceed if lag is within bounds - lag = self.collect_train_lag - return lag < self.off_policy_degree - - async def run_collect( - self, - collect_fn: Callable[[], Awaitable[Any]], - ) -> Any: - """ - Run collection with coordination. - - Args: - collect_fn: Async collection function - - Returns: - Collection result - """ - # Wait if needed (for sync mode or if lag is too high) - while not self.can_collect(): - logger.debug(f"Collect waiting (lag={self.collect_train_lag}, limit={self.off_policy_degree})") - await asyncio.sleep(0.1) - - # Run collection - start_time = time.time() - result = await collect_fn() - elapsed = time.time() - start_time - - # Update counter - async with self._lock: - self._collect_count += 1 - self._collect_times.append(elapsed) - - logger.debug(f"Collect completed in {elapsed:.2f}s (count={self._collect_count})") - return result - - async def run_train( - self, - train_fn: Callable[[], Awaitable[Any]], - ) -> Any: - """ - Run training with coordination. - - Args: - train_fn: Async training function - - Returns: - Training result - """ - # Wait if needed - while not self.can_train(): - logger.debug(f"Train waiting (collect={self._collect_count}, train={self._train_count})") - await asyncio.sleep(0.1) - - # Run training - start_time = time.time() - result = await train_fn() - elapsed = time.time() - start_time - - # Update counter - async with self._lock: - self._train_count += 1 - self._train_times.append(elapsed) - - logger.debug(f"Train completed in {elapsed:.2f}s (count={self._train_count}, lag={self.collect_train_lag})") - return result - - async def run_eval( - self, - eval_fn: Callable[[], Awaitable[Any]], - ) -> Any: - """ - Run evaluation with coordination. - - Args: - eval_fn: Async evaluation function - - Returns: - Evaluation result - """ - start_time = time.time() - - if self.enable_async_eval: - # Cancel previous eval if still running - if self._eval_task is not None and not self._eval_task.done(): - logger.info("Cancelling previous eval task") - self._eval_task.cancel() - try: - await self._eval_task - except asyncio.CancelledError: - pass - - # Run eval in background - self._eval_task = asyncio.create_task(eval_fn()) - logger.info("Started async eval in background") - - # Return immediately (don't wait) - return None - else: - # Synchronous eval - result = await eval_fn() - elapsed = time.time() - start_time - self._eval_times.append(elapsed) - logger.debug(f"Eval completed in {elapsed:.2f}s") - return result - - async def wait_for_eval(self) -> Optional[Any]: - """ - Wait for async eval to complete (if running). - - Returns: - Eval result if eval was running, None otherwise - """ - if self._eval_task is not None and not self._eval_task.done(): - logger.info("Waiting for async eval to complete...") - try: - result = await self._eval_task - return result - except asyncio.CancelledError: - logger.warning("Eval task was cancelled") - return None - return None - - def get_statistics(self) -> Dict[str, Any]: - """ - Get performance statistics. - - Returns: - Dictionary with timing statistics - """ - stats = { - 'collect_count': self._collect_count, - 'train_count': self._train_count, - 'collect_train_lag': self.collect_train_lag, - 'mode': 'synchronous' if self.is_synchronous else 'asynchronous', - } - - if self._collect_times: - stats['collect_avg_time'] = sum(self._collect_times) / len(self._collect_times) - stats['collect_total_time'] = sum(self._collect_times) - - if self._train_times: - stats['train_avg_time'] = sum(self._train_times) / len(self._train_times) - stats['train_total_time'] = sum(self._train_times) - - if self._eval_times: - stats['eval_avg_time'] = sum(self._eval_times) / len(self._eval_times) - stats['eval_total_time'] = sum(self._eval_times) - - return stats - - def reset_counters(self): - """Reset all counters (useful for testing).""" - self._collect_count = 0 - self._train_count = 0 - self._collect_times.clear() - self._train_times.clear() - self._eval_times.clear() - logger.info("AsyncTrainingCoordinator counters reset") - - -async def run_async_training_loop( - coordinator: AsyncTrainingCoordinator, - collect_fn: Callable[[], Awaitable[Any]], - train_fn: Callable[[], Awaitable[Any]], - eval_fn: Callable[[], Awaitable[Any]], - eval_interval: int, - max_iterations: int, -): - """ - Main async training loop that coordinates collect/train/eval. - - Args: - coordinator: AsyncTrainingCoordinator instance - collect_fn: Async collection function - train_fn: Async training function - eval_fn: Async evaluation function - eval_interval: How often to run eval (in iterations) - max_iterations: Maximum training iterations - """ - logger.info(f"Starting async training loop (max_iter={max_iterations})") - - if coordinator.is_synchronous: - # ======================================================================== - # SYNCHRONOUS MODE: Original serial execution - # ======================================================================== - logger.info("Running in SYNCHRONOUS mode") - - for iteration in range(max_iterations): - # 1. Collect - logger.info(f"[Iter {iteration}] Collecting...") - await coordinator.run_collect(collect_fn) - - # 2. Train - logger.info(f"[Iter {iteration}] Training...") - await coordinator.run_train(train_fn) - - # 3. Eval (if needed) - if iteration % eval_interval == 0: - logger.info(f"[Iter {iteration}] Evaluating...") - await coordinator.run_eval(eval_fn) - - else: - # ======================================================================== - # ASYNCHRONOUS MODE: Concurrent execution with bounded lag - # ======================================================================== - logger.info(f"Running in ASYNCHRONOUS mode (off_policy_degree={coordinator.off_policy_degree})") - - # Create tasks for collect and train - collect_task = None - train_tasks = [] - - iteration = 0 - while iteration < max_iterations: - tasks_to_wait = [] - - # Start collect if allowed - if coordinator.can_collect() and (collect_task is None or collect_task.done()): - logger.debug(f"[Iter {iteration}] Starting collect task") - collect_task = asyncio.create_task(coordinator.run_collect(collect_fn)) - tasks_to_wait.append(collect_task) - - # Start train if allowed and there's data - if coordinator.can_train(): - logger.debug(f"[Iter {iteration}] Starting train task") - train_task = asyncio.create_task(coordinator.run_train(train_fn)) - train_tasks.append(train_task) - tasks_to_wait.append(train_task) - iteration += 1 - - # Eval (if needed) - if iteration % eval_interval == 0 and iteration > 0: - logger.info(f"[Iter {iteration}] Triggering eval") - await coordinator.run_eval(eval_fn) - - # Wait for at least one task to complete - if tasks_to_wait: - done, pending = await asyncio.wait(tasks_to_wait, return_when=asyncio.FIRST_COMPLETED) - logger.debug(f"Tasks completed: {len(done)}, pending: {len(pending)}") - else: - # No tasks ready, wait a bit - await asyncio.sleep(0.1) - - # Clean up completed train tasks - train_tasks = [t for t in train_tasks if not t.done()] - - # Wait for all remaining tasks - logger.info("Waiting for remaining tasks to complete...") - if collect_task and not collect_task.done(): - await collect_task - for task in train_tasks: - if not task.done(): - await task - - # Wait for eval if running - await coordinator.wait_for_eval() - - # Print statistics - stats = coordinator.get_statistics() - logger.info("="*80) - logger.info("Training Loop Statistics:") - logger.info(f" Mode: {stats['mode']}") - logger.info(f" Collect count: {stats['collect_count']}") - logger.info(f" Train count: {stats['train_count']}") - logger.info(f" Final lag: {stats['collect_train_lag']}") - if 'collect_avg_time' in stats: - logger.info(f" Avg collect time: {stats['collect_avg_time']:.2f}s") - if 'train_avg_time' in stats: - logger.info(f" Avg train time: {stats['train_avg_time']:.2f}s") - if 'eval_avg_time' in stats: - logger.info(f" Avg eval time: {stats['eval_avg_time']:.2f}s") - logger.info("="*80) diff --git a/zoo/jericho/priorzero/atari_action_meanings.py b/zoo/jericho/priorzero/atari_action_meanings.py new file mode 100644 index 000000000..29dad5540 --- /dev/null +++ b/zoo/jericho/priorzero/atari_action_meanings.py @@ -0,0 +1,187 @@ +""" +Atari Action Space Mapping + +Maps integer action indices to semantic action names for better VL understanding. +""" + +# Atari action space mappings +# Source: https://github.com/openai/gym/blob/master/gym/envs/atari/atari_env.py +ATARI_ACTION_MEANINGS = { + 'PongNoFrameskip-v4': { + 0: 'NOOP', + 1: 'FIRE', + 2: 'RIGHT', + 3: 'LEFT', + 4: 'RIGHTFIRE', + 5: 'LEFTFIRE', + }, + 'BreakoutNoFrameskip-v4': { + 0: 'NOOP', + 1: 'FIRE', + 2: 'RIGHT', + 3: 'LEFT', + }, + 'SpaceInvadersNoFrameskip-v4': { + 0: 'NOOP', + 1: 'FIRE', + 2: 'RIGHT', + 3: 'LEFT', + 4: 'RIGHTFIRE', + 5: 'LEFTFIRE', + }, + 'QbertNoFrameskip-v4': { + 0: 'NOOP', + 1: 'FIRE', + 2: 'UP', + 3: 'RIGHT', + 4: 'LEFT', + 5: 'DOWN', + }, + 'MsPacmanNoFrameskip-v4': { + 0: 'NOOP', + 1: 'UP', + 2: 'RIGHT', + 3: 'LEFT', + 4: 'DOWN', + 5: 'UPRIGHT', + 6: 'UPLEFT', + 7: 'DOWNRIGHT', + 8: 'DOWNLEFT', + }, + 'SeaquestNoFrameskip-v4': { + 0: 'NOOP', + 1: 'FIRE', + 2: 'UP', + 3: 'RIGHT', + 4: 'LEFT', + 5: 'DOWN', + 6: 'UPRIGHT', + 7: 'UPLEFT', + 8: 'DOWNRIGHT', + 9: 'DOWNLEFT', + 10: 'UPFIRE', + 11: 'RIGHTFIRE', + 12: 'LEFTFIRE', + 13: 'DOWNFIRE', + 14: 'UPRIGHTFIRE', + 15: 'UPLEFTFIRE', + 16: 'DOWNRIGHTFIRE', + 17: 'DOWNLEFTFIRE', + }, + 'MontezumaRevengeNoFrameskip-v4': { + 0: 'NOOP', + 1: 'FIRE', + 2: 'UP', + 3: 'RIGHT', + 4: 'LEFT', + 5: 'DOWN', + 6: 'UPRIGHT', + 7: 'UPLEFT', + 8: 'DOWNRIGHT', + 9: 'DOWNLEFT', + 10: 'UPFIRE', + 11: 'RIGHTFIRE', + 12: 'LEFTFIRE', + 13: 'DOWNFIRE', + 14: 'UPRIGHTFIRE', + 15: 'UPLEFTFIRE', + 16: 'DOWNRIGHTFIRE', + 17: 'DOWNLEFTFIRE', + }, + 'LunarLander-v2': { + 0: 'NOOP', + 1: 'LEFT_ENGINE', + 2: 'MAIN_ENGINE', + 3: 'RIGHT_ENGINE', + }, +} + + +def get_action_meanings(env_id: str, action_space_size: int) -> dict: + """ + Get action meanings for a given Atari environment. + + Args: + env_id: Environment ID (e.g., 'PongNoFrameskip-v4') + action_space_size: Number of actions in the action space + + Returns: + Dictionary mapping action indices to semantic names + """ + if env_id in ATARI_ACTION_MEANINGS: + return ATARI_ACTION_MEANINGS[env_id] + + # Fallback: generic action names + return {i: f'ACTION_{i}' for i in range(action_space_size)} + + +def action_index_to_name(env_id: str, action_index: int, action_space_size: int) -> str: + """ + Convert action index to semantic name. + + Args: + env_id: Environment ID + action_index: Action index (0, 1, 2, ...) + action_space_size: Total number of actions + + Returns: + Semantic action name (e.g., 'FIRE', 'RIGHT') + """ + meanings = get_action_meanings(env_id, action_space_size) + return meanings.get(action_index, f'ACTION_{action_index}') + + +def action_name_to_index(env_id: str, action_name: str, action_space_size: int) -> int: + """ + Convert semantic action name to index. + + Args: + env_id: Environment ID + action_name: Semantic action name (e.g., 'FIRE', 'RIGHT') + action_space_size: Total number of actions + + Returns: + Action index (0, 1, 2, ...) + """ + meanings = get_action_meanings(env_id, action_space_size) + + # Create reverse mapping + name_to_idx = {name: idx for idx, name in meanings.items()} + + # Try exact match first + if action_name in name_to_idx: + return name_to_idx[action_name] + + # Try case-insensitive match + action_name_upper = action_name.upper() + if action_name_upper in name_to_idx: + return name_to_idx[action_name_upper] + + # Try parsing "ACTION_X" format + if action_name.startswith('ACTION_'): + try: + return int(action_name.split('_')[1]) + except (IndexError, ValueError): + pass + + # Fallback: return 0 (NOOP) + return 0 + + +if __name__ == '__main__': + # Test + print("Testing Atari action mappings:") + print("\nPong actions:") + for i in range(6): + name = action_index_to_name('PongNoFrameskip-v4', i, 6) + print(f" {i} -> {name}") + + print("\nBreakout actions:") + for i in range(4): + name = action_index_to_name('BreakoutNoFrameskip-v4', i, 4) + print(f" {i} -> {name}") + + print("\nReverse mapping (Pong):") + for name in ['NOOP', 'FIRE', 'RIGHT', 'LEFT']: + idx = action_name_to_index('PongNoFrameskip-v4', name, 6) + print(f" {name} -> {idx}") diff --git a/zoo/jericho/priorzero/docs/priorzero_core_mechanism_debug_analysis_20260322.md b/zoo/jericho/priorzero/docs/priorzero_core_mechanism_debug_analysis_20260322.md new file mode 100644 index 000000000..94b3bc4b9 --- /dev/null +++ b/zoo/jericho/priorzero/docs/priorzero_core_mechanism_debug_analysis_20260322.md @@ -0,0 +1,577 @@ +# PriorZero 核心机制与调试分析文档 + +> 生成时间: 2026-03-22 | 基于最新代码(含 PR #441 重构,VL 模型支持、rollout_logprob 重命名、样本去重、扩展训练指标) + +--- + +## 1. PriorZero 整体架构与数据流 + +### 1.1 Actor-Critic 交互流程 + +PriorZero 在 Jericho 文本冒险环境下,采用 **WM-LLM (World Model) + Policy LLM** 双模型协同架构: + +- **WM-LLM (World Model)**:基于 UniZero 的 transformer-based world model,负责环境建模、value 预测、policy logits 生成 +- **Policy LLM**:基于 Qwen2.5 系列(含 VL 变体)的因果语言模型,通过 PPO/GSPO 进行策略优化,输出动作的 token-level log-probability。最新代码通过 `AutoConfig` 自动检测 VL 模型并使用 `AutoModelForVision2Seq`(`actor.py:92-113`) + +两者通过 **交替训练 (alternating training)** 机制协调:先训练 WM 若干轮,再训练 LLM 若干轮,循环往复。 + +### 1.2 核心数据流 + +``` +┌─────────────────────────────────────────────────────────────────────┐ +│ ROLLOUT (Rank 0) │ +│ │ +│ Jericho Env ──→ raw_obs_text, valid_actions, history │ +│ │ │ +│ ▼ │ +│ DataProcessor.get_llm_prior() │ +│ ├─ [可选] _build_cot_prefix_texts() → CoT reasoning prefix │ +│ ├─ _score_labels_with_prompt_logprobs() → per-action logprob │ +│ │ (vLLM prompt_logprobs=1, 拼接 context+label 后提取) │ +│ └─ 返回: llm_prior_per_seq, llm_prior_per_tok, cot_prefixes │ +│ │ (tok_dict 中 key 为 'rollout_action_logprob') │ +│ ▼ │ +│ Policy._forward_collect(llm_prior_logprob=...) │ +│ ├─ WM initial_inference() → wm_policy_logits, wm_value │ +│ ├─ 融合 LLM + WM logits (fixed/adaptive 加权) │ +│ └─ MCTS search → 选择动作 │ +│ │ │ +│ ▼ │ +│ GameSegment.append(raw_obs, history, llm_prior_per_tok, │ +│ cot_prefix, llm_action) │ +└───────────────────────┬─────────────────────────────────────────────┘ + │ + ▼ +┌─────────────────────────────────────────────────────────────────────┐ +│ BUFFER (Rank 0) │ +│ │ +│ PriorZeroGameBufferOptimized │ +│ ├─ push_game_segments(new_data) │ +│ ├─ sample(batch_size) → WM 训练数据 │ +│ └─ fetch_latest_batch() → LLM 训练数据 (priorzero_batch) │ +│ 返回: (raw_obs_list, history_obs_list, │ +│ llm_prior_per_tok_list, target_value, │ +│ pred_value, cot_prefix_list, llm_action_list) │ +└───────────────────────┬─────────────────────────────────────────────┘ + │ + ▼ +┌─────────────────────────────────────────────────────────────────────┐ +│ TRAIN (All Ranks via broadcast) │ +│ │ +│ ── WM Phase ── │ +│ learner.train(train_data) → WM losses (obs, reward, policy, value) │ +│ │ +│ ── LLM Phase ── │ +│ 1. bcast_obj(priorzero_batch) → 广播到所有 rank │ +│ 2. DataProcessor.make_llm_train_samples() │ +│ ├─ build_llm_samples() → advantage = target_value - pred_value │ +│ ├─ unique_dicts_hash() 去重(datafactory.py:46-58) │ +│ ├─ advantage normalization (batch_norm / running_norm) │ +│ ├─ [可选] format_reward 融合 │ +│ └─ tokenize + pad → 返回 (flag, (input_ids, attn_mask, │ +│ action_mask, advantage, rollout_logprob, log_status)) │ +│ 3. PriorZeroLLMTrainer.train_batch() │ +│ ├─ PolicyModel.forward() → old_action_log_probs (当前策略) │ +│ ├─ [可选] ReferenceModel.forward() → ref_log_probs │ +│ ├─ PolicyModel.fit(batch_data, kl_ctl) │ +│ │ └─ BatchPPOTrainer.train_batch() │ +│ │ ├─ Actor.forward() → action_log_probs │ +│ │ ├─ PolicyLoss(log_probs, old_log_probs, advantages, │ +│ │ │ action_mask, rollout_log_probs) │ +│ │ │ → actor_loss, clipfrac, approx_kl, vllm_kl │ +│ │ ├─ KL loss (vs reference model) │ +│ │ ├─ Entropy loss │ +│ │ ├─ 扩展指标: ratio_mean/std, adv_mean/std, │ +│ │ │ log_prob_new/old_mean, kl_coef, total_loss │ +│ │ └─ backward + optimizer_step │ +│ └─ broadcast_to_vllm() → 同步权重到 vLLM engine │ +└─────────────────────────────────────────────────────────────────────┘ +``` + +### 1.3 WM-LLM 角色流转 + +| 阶段 | WM 角色 | Policy LLM 角色 | +|------|---------|-----------------| +| WM Warmup | 训练中(obs/reward/policy/value loss) | 冻结,仅用 vLLM 提供 prior | +| WM Phase | 训练中 | 冻结,仅用 vLLM 提供 prior | +| LLM Phase | 冻结(提供 target_value, pred_value) | 训练中(PPO/GSPO loss) | +| Collect | 推理(initial_inference) | 推理(vLLM 计算 prior) | + +关键控制逻辑在 `priorzero_entry_sync.py:242-294`: +```python +# WM phase +if llm_cfg.enable_world_model and current_phase == "wm": + for i in range(update_per_collect): + train_data = replay_buffer.sample(batch_size, policy) + learner.train(train_data) + if learner.train_iter - last_wm_train_iter >= train_schedule["wm_update_iters"]: + current_phase = "llm" + +# LLM phase +if llm_cfg.enable_rft and current_phase == "llm": + priorzero_batch = replay_buffer.fetch_latest_batch(batch_size=-1, policy=policy) + torch.cuda.empty_cache() # 清理 policy 的 cache,防止 OOM(entry_sync.py:264) + flag, train_samples = data_processor.make_llm_train_samples(priorzero_batch, ...) + if not flag: # 样本不足时跳过(entry_sync.py:282) + continue + trainer.train_batch(train_samples) + if trainer.global_step - last_llm_train_iter >= train_schedule["llm_update_iters"]: + current_phase = "wm" + data_processor.value_normalizer.clear() # ← 关键:切换回 WM 时清空 normalizer +``` + +--- + +## 2. 核心机制深度解析 + +### 2.1 Value Normalization (`value_running_norm`) + +**代码位置**:`models/stability_optimizer.py:AdaptiveValueNormalizer`,调用点在 `priorzero_datafactory.py:368-440` + +#### 更新逻辑 + +1. **输入**:`advantage = target_value - pred_value`(由 WM 提供的 TD bootstrap value 差) +2. **裁剪**(可选): + - `soft`:`f(x) = sign(x) * log(1 + |x|)` → 压缩极端值但保留符号(`stability_optimizer.py:47-51`) + - `hard`:分位数裁剪,保留 `[2.5%, 97.5%]` 区间(`stability_optimizer.py:53-66`) +3. **EMA 统计量更新**: + ```python + # stability_optimizer.py:102-108 + momentum = init_momentum + (final_momentum - init_momentum) * min(update_count / warmup_steps, 1.0) + if update_count == 0: + running_mean = batch_mean + running_std = batch_std + else: + running_mean = momentum * running_mean + (1 - momentum) * batch_mean + running_std = momentum * running_std + (1 - momentum) * batch_std + ``` +4. **归一化**:`y = (x - running_mean) / (running_std + 1e-6)`(`stability_optimizer.py:114`) + +#### 清空机制及其影响 + +**清空时机**:`priorzero_entry_sync.py:293-294` +```python +if train_alternate and trainer.global_step - last_llm_train_iter >= train_schedule["llm_update_iters"]: + current_phase = "wm" + if data_processor.value_normalizer is not None: + data_processor.value_normalizer.clear() # reset running_mean=0, running_std=1, update_count=0 +``` + +**`clear()` 实现**(`stability_optimizer.py:146-155`): +```python +def clear(self): + self.running_mean = 0.0 + self.running_std = 1.0 + self.update_count = 0 + self.value_history.clear() +``` + +**影响分析**: +- **正面**:每轮 WM 训练后 value 分布可能显著变化(WM 学到了新东西),清空 EMA 防止旧统计量产生 stale bias +- **负面风险**:清空后第一个 batch 的 `update_count=0`,直接用 batch 统计量初始化 running stats。若该 batch 恰好含极端值,会导致归一化后的 advantage 尺度不稳定 +- **调试建议**:监控每次 `clear()` 后首个 batch 的 `norm_min/norm_max`,若出现极端值(>10 或 <-10),考虑在 clear 后保留 `running_std` 的下界 + +### 2.2 Advantage 计算与截断 + +**代码位置**:`priorzero_datafactory.py:345-442` + +#### 计算方式 + +**非 GAE**,而是直接的 TD-error: +```python +# priorzero_datafactory.py:349 +advantage = target_value - pred_value +``` +其中: +- `target_value[t]`:从时刻 t 开始的 `td_step` 步真实奖励折扣和 + bootstrap `V(t + td_step)` +- `pred_value[t]`:WM 在时刻 t 的 value 预测 `V(t)` + +#### 三种归一化模式 + +| 模式 | 代码位置 | 说明 | +|------|---------|------| +| `advantage` | `datafactory.py:351-356` | 原始值不变,最简单但尺度不可控 | +| `advantage_batch_norm` | `datafactory.py:359-366` | `(adv - mean) / (std + 1e-8)` 当前 batch 归一化 | +| `advantage_running_norm` | `datafactory.py:368-440` | `AdaptiveValueNormalizer`(EMA + soft/hard clip)或 fallback 手动 EMA | + +#### 截断处理 + +- **soft clip**:`sign(x) * log(1 + |x|)`,阈值判定 `|x| > 10` 时计入 `clipped_count` +- **hard clip**:分位数 `[2.5%, 97.5%]`,前 `hard_clip_start_updates=10` 次不启用 +- **注意**:截断在归一化 **之前** 应用,先压缩极端值再计算统计量 + +#### Format Reward 融合(可选) + +```python +# priorzero_datafactory.py:354-356 +if fmt_rewards is not None: + advantage = (1 - fmt_weight) * advantage + fmt_weight * fmt_rewards +``` +`fmt_rewards` 为 0/1 二值,检查输出是否符合 `Reasoning: ...\nAction: ...` 格式(`_format_reward()` at `datafactory.py:17-42`)。 + +#### 样本去重机制 + +最新代码在 `make_llm_train_samples()` 中增加了基于 hash 的去重: +```python +# datafactory.py:294-304 +if len(samples) >= max_samples: + unique_samples = unique_dicts_hash(samples) # MD5 hash 去重 + if len(unique_samples) >= max_samples: + samples = unique_samples[:max_samples] + else: + remain = max_samples - len(unique_samples) + samples = unique_samples + samples[:remain] # 不足时用重复样本补齐 +else: + return False, samples # 样本不足,返回 flag=False +``` + +`unique_dicts_hash()`(`datafactory.py:46-58`)通过 `pickle.dumps` + `md5` 对每个样本 dict 做去重。 + +### 2.3 Importance Sampling (IS) 与 `clipfrac` + +**代码位置**:`models/loss.py:PolicyLoss` + +#### IS Ratio 计算 + +当前代码中存在 **三层 logprob**,理解其区别至关重要: + +| 变量名 | 来源 | 含义 | +|--------|------|------| +| `rollout_action_logprob` | vLLM 在 collect 时计算 | `π_rollout(a|s)` — rollout 策略的 logprob | +| `old_action_log_probs` | `PolicyModel.forward()` 在 LLM 训练开始前计算 | `π_θ_old(a|s)` — 当前 epoch 开始时的策略 | +| `action_log_probs` | `Actor.forward()` 在每个 micro-batch 中计算 | `π_θ(a|s)` — 正在更新中的策略 | + +PPO 标准 ratio(`loss.py:52-54`): +```python +log_ratio = log_probs - old_log_probs # π_θ / π_θ_old(同一 epoch 内的变化) +ratio = log_ratio.exp() +``` + +vLLM IS correction ratio(`loss.py:84-97`,仅 `enable_vllm_is_correction=True` 时): +```python +vllm_is = exp(old_log_probs - rollout_log_probs) # π_θ_old / π_rollout(跨 epoch 的偏移) +vllm_is = vllm_is.clamp(low_threshold, high_threshold) +loss = vllm_is * loss # 修正 off-policy 偏差 +``` + +#### PPO Clipped Surrogate Loss + +```python +# loss.py:68-73 +surr1 = ratio * advantages +surr2 = ratio.clamp(1 - eps_low, 1 + eps_high) * advantages # 默认 [0.8, 1.2] +loss = -torch.min(surr1, surr2) +``` + +Dual-clip 变体(`loss.py:74-80`):当 advantage < 0 时额外增加下界 `dual_clip * advantages`。 + +ICEPOP 变体(`loss.py:86-90`):区间外的 IS 权重直接置零(而非 clamp)。 + +#### `clipfrac` 指标含义 + +```python +# loss.py:104-105 +clipped = ratio.gt(1 + eps_high) | ratio.lt(1 - eps_low) +clipfrac = masked_mean(clipped, action_mask, dim=None) +``` + +- **含义**:token 级别的 IS ratio 落在 `[1-ε, 1+ε]` 区间 **之外** 的比例 +- **健康值**:`clipfrac ∈ [0.05, 0.3]` + - `< 0.05`:策略更新太保守,学习效率低 + - `> 0.5`:策略偏移严重,PPO clip 大量生效,可能导致训练不稳定 +- **相关指标**:`clip_ratio = P(surr2 < surr1)` 表示 clip 实际约束了多少 loss + +#### `approx_kl` 计算 + +```python +# loss.py:108 +approx_kl = masked_mean(-log_ratio.detach(), action_mask, dim=None) +``` +即 `E[-log(π_θ/π_old)] ≈ KL(π_old || π_θ)`,Schulman k1 近似。 + +#### 新增训练指标(`actor.py:324-365`) + +最新代码在 `BatchPPOTrainer.train_batch()` 中新增了以下诊断指标: + +| 指标 | 计算方式 | 诊断价值 | +|------|---------|---------| +| `ratio_mean` | `masked_mean(exp(log_probs - old_log_probs))` | IS ratio 均值,健康值 ≈ 1.0 | +| `ratio_std` | IS ratio 的标准差 | 偏移幅度,过大说明策略变化剧烈 | +| `advantage_mean/std` | 当前 micro-batch 的 advantage 统计 | 监控 advantage 分布 | +| `log_prob_new_mean` | 当前策略 log_prob 均值 | 策略信心度 | +| `log_prob_old_mean` | 旧策略 log_prob 均值 | 基线参考 | +| `total_loss` | `actor_loss + kl_loss * kl_coef - entropy * entropy_coef` | 含所有正则项的完整 loss | +| `kl_coef` | `float(kl_ctl.value)` | 当前 KL penalty 系数 | +| `vllm_kl` | `masked_mean(rollout_logprobs - old_logprobs)` | vLLM IS 校正时的 KL 散度 | + +### 2.4 异步采样控制 (`max_rollout_staleness`) + +**代码位置**:`priorzero_entry_sync.py:280` + +#### 控制逻辑 + +```python +llm_need_sample_cnt = llm_cfg.train_batch_size * llm_cfg.max_rollout_staleness // 1 +flag, train_samples = data_processor.make_llm_train_samples(priorzero_batch, max_samples=llm_need_sample_cnt) +``` + +- `max_rollout_staleness` 控制 LLM 训练时允许使用多少倍于 `train_batch_size` 的样本 +- 默认值 `1`:只用最近一次 collect 的数据量(= `train_batch_size` 个样本) +- 值越大,允许使用越多"旧"数据,提升样本效率但增加 off-policy 程度 + +#### 返回值变化 + +最新代码中 `make_llm_train_samples()` 返回 `(flag, data)` 元组(`datafactory.py:304, 462`): +- `flag=True`:样本充足,`data` 为训练数据元组 +- `flag=False`:样本不足(`< max_samples`),`data` 为原始 samples 列表(非训练格式) + +调用方通过 `if not flag: continue` 跳过本轮 LLM 训练(`entry_sync.py:282-284`)。 + +#### 过时轨迹处理 + +当前实现中,过时数据不是通过时间戳丢弃的,而是通过 **buffer 的 `mark_latest_transitions_consumed()` + `fetch_latest_batch()`** 机制: + +```python +# priorzero_entry_sync.py:287 +replay_buffer.mark_latest_transitions_consumed() # 标记当前数据已消费 +``` + +`fetch_latest_batch(batch_size=-1)` 只返回自上次 `mark` 以来新增的数据。因此 `max_rollout_staleness` 实际控制的是**每次 LLM 训练使用的样本上限**,而非数据的"年龄"。 + +--- + +## 3. 三大 Bug/痛点排查指南 + +### 3.1 痛点一:Policy Loss 出现 NaN + +**现象**:LLM 接近最优时 KL 变大,固定 LR 下后期 Loss 变 NaN。 + +#### 潜在原因 1:Value Normalizer 清空后首 batch 极端值 + +**风险点**:`priorzero_entry_sync.py:293-294` 调用 `value_normalizer.clear()` 后: +- `update_count` 重置为 0 +- 首 batch 直接赋值 `running_mean = batch_mean, running_std = batch_std` +- 若 WM 刚训练完 value 分布剧变,首 batch 可能包含极端 advantage +- `stability_optimizer.py:114`: `y = (x - running_mean) / (running_std + 1e-6)` — 若 `running_std` 极小(batch 中所有 advantage 接近),归一化后值可能爆炸 + +**修复建议**: +```python +# 在 clear() 中保留 std 下界 +def clear(self): + self.running_mean = 0.0 + self.running_std = max(1.0, self.running_std * 0.5) # 不完全重置 std + self.update_count = 0 + self.value_history.clear() +``` + +#### 潜在原因 2:KL 散度计算中的数值溢出 + +**风险点**:`utils.py:60-94` 中的 `compute_approx_kl()` + +```python +# k3 estimator (utils.py:88-91) +log_ratio = log_probs - log_probs_base # 当策略偏移很大时,可能是很大的正/负数 +log_ratio = -log_ratio +log_ratio = log_ratio.exp() - 1 - log_ratio # exp(大正数) → Inf → NaN +``` + +虽然有 `log_ratio.clamp(min=-10, max=10)`(line 93),但 clamp 在 **最后** 应用,此时 `exp()` 可能已经溢出。 + +**修复建议**:将 clamp 移到 `exp()` 之前: +```python +if kl_estimator == "k3": + log_ratio = log_probs.float() - log_probs_base.float() + log_ratio = (-log_ratio).clamp(min=-10, max=10) # 先 clamp 再 exp + log_ratio = log_ratio.exp() - 1 + (log_probs.float() - log_probs_base.float()) +``` + +#### 潜在原因 3:log_probs 在 bfloat16 下精度不足 + +**风险点**:`actor.py:150` +```python +output["logits"] = output["logits"].to(torch.float32) +``` +虽然 logits 转了 float32,但 `log_probs_from_logits()` 中的 `flash_attn cross_entropy_loss` 路径(`utils.py:121`)可能在内部回退到低精度。 + +**排查方法**:利用最新代码中的扩展指标,在 `BatchPPOTrainer.train_batch()` 的 `actor.py:324-339` 处已自动记录 `ratio_mean/std`、`log_prob_new/old_mean`。观察这些指标是否出现 NaN/Inf 前兆: +```python +# 已有的指标(无需额外添加代码) +# ratio_mean ≈ 1.0 是健康的;>> 1 或 << 1 说明策略偏移严重 +# log_prob_new_mean 与 log_prob_old_mean 的差值 ≈ approx_kl +``` + +若需更细粒度排查,可添加: +```python +# 在 actor_loss 计算后添加(actor.py:286 之后) +if torch.isnan(actor_loss) or torch.isinf(actor_loss): + print(f"[NaN DEBUG] action_log_probs: min={action_log_probs.min()}, max={action_log_probs.max()}") + print(f"[NaN DEBUG] old_log_probs: min={micro_batch['old_action_log_probs'].min()}, max={micro_batch['old_action_log_probs'].max()}") + print(f"[NaN DEBUG] advantages: min={micro_batch['advantages'].min()}, max={micro_batch['advantages'].max()}") +``` + +#### 潜在原因 4:Advantage 极端值未被充分抑制 + +当 `advantage_type="advantage"`(无归一化)时,raw advantage 可能非常大。PPO ratio * advantage 的乘积导致梯度爆炸。 + +**排查**:监控 `value_advantage_max/min` 和新增的 `advantage_mean/std` 指标,若 `|adv| > 100` 需要启用 `advantage_running_norm`。 + +### 3.2 痛点二:LLM 与 vLLM 的 `logprob` 差异 + +**现象**:相同输入输出下,原生 LLM(`Actor.forward()`)和 vLLM(`_score_labels_with_prompt_logprobs()`)给出的 logprob 差异很大。 + +#### 差异来源 1:Temperature 处理不一致 + +- **vLLM 侧**:`priorzero_datafactory.py:609-611` + ```python + sampling_params = SamplingParams(temperature=self.temperature, ...) + ``` + vLLM 的 `prompt_logprobs` 返回的是 **经过 temperature 缩放后** 的 logprob(`logit / T` 后做 log_softmax) + +- **Actor 侧**:`actor.py:157` + ```python + log_probs = log_probs_from_logits(output["logits"], rolled_sequences, temperature=self.temperature) + ``` + `utils.py:112-113`:`logits.div_(temperature)` **原地修改**后做 log_softmax + +- **风险**:如果两侧的 `temperature` 配置不一致(`llm_cfg.temperature` vs `strategy.args.temperature`),logprob 会系统性偏移。**特别注意**:`Actor.__init__` 的 `temperature` 来自 `strategy.args.temperature`(`actor.py:551`),`DataProcessor` 的 `self.temperature` 也来自 `strategy.args.temperature`(`datafactory.py:84`),理论上应一致,但需确认。 + +**排查方法**: +```python +# 在 train_batch 中对比 +print(f"Actor temperature: {self.actor.temperature}") +print(f"vLLM SamplingParams temperature: {data_processor.temperature}") +``` + +#### 差异来源 2:Tokenization 对齐问题 + +- **vLLM 侧**:`priorzero_datafactory.py:618-636` + ```python + context_ids = tokenizer(all_context_texts, add_special_tokens=False, ...)["input_ids"] + label_ids = tokenizer(label_texts, add_special_tokens=False, ...)["input_ids"] + full_ids = [c + l for c, l in zip(context_ids, label_ids)] # 手动拼接 + ``` + 然后通过 `prompt_token_ids=full_ids` 传给 vLLM + +- **Actor 侧**:`priorzero_datafactory.py:329` + ```python + inputs = self.tokenizer.pad({"input_ids": full_ids_list}, padding=True, return_tensors="pt") + ``` + 使用同一套 `full_ids`(从 sample 中取出的),通过 **左填充 (padding_side="left")** 对齐 + +- **关键风险**:**padding 引入的 attention_mask 差异**。vLLM 不做 padding,直接处理变长序列;Actor 做左 padding 但依赖 `attention_mask` 和 `position_ids` 正确排除 pad tokens。若 `position_ids` 计算有误(`actor.py:146-147`),会导致 logprob 偏移。 + +#### 差异来源 3:BOS Token 处理 + +- **vLLM**:`prompt_logprobs[0]` 是 `None`(第一个 token 无条件概率),从 `j=1` 开始提取(`datafactory.py:651`) +- **Actor**:`log_probs = log_probs[:, :-1]`(`actor.py:159`),即 logits 右移一位后取 log_softmax + +两侧都跳过了第一个 token,理论一致。但如果 `apply_chat_template()` 在 vLLM 和 Actor 侧产生不同的 BOS/前缀 token,会导致 context 长度不同。 + +**排查方法**:利用新增的 `log_prob_new_mean` 和 `log_prob_old_mean` 指标,对比两者与 `rollout_action_logprob` 的差异: +```python +# 在 train_batch 开始时对比 token-level logprob +actor_lp = action_log_probs[0] # [T_action] +vllm_lp = micro_batch['rollout_action_logprob'][0] # [T_action] +mask = micro_batch['action_mask'][0] +print(f"Actor logprob (masked): {(actor_lp * mask).sum()}") +print(f"vLLM logprob (masked): {(vllm_lp * mask).sum()}") +print(f"Diff per token: {((actor_lp - vllm_lp) * mask).abs().max()}") +``` + +#### 差异来源 4:vLLM model_impl 与 HF 实现差异 + +`vllm_engine.py` 配置 `model_impl="transformers"`,理论上使用与 HF 相同的模型实现。但 vLLM 的 attention kernel(即使用 eager 模式)、数值精度路径可能与 HF + flash_attention_2 存在微小差异。 + +**注意**:最新代码中 Actor 新增了 VL 模型支持(`actor.py:92-113`),若使用 VL 模型,vLLM 侧也需确保使用对应的 VL 推理路径。 + +### 3.3 痛点三:CoT (Chain of Thought) 融合优化 + +**需求**:在无 CoT 的最佳 config 基础上,加入 `weight=0.1` 的 CoT loss。 + +#### 当前 CoT 实现分析 + +**CoT 生成**:`priorzero_datafactory.py:464-516` (`_build_cot_prefix_texts()`) +- 使用 vLLM 生成 "Reasoning: ... \nAction:" 格式的推理前缀 +- Stop condition: `"\n\n"` +- 生成后截取到 "Action:" 标记处 + +**CoT 融入训练**:`priorzero_datafactory.py:316-324` +```python +if self.use_cot: + targets_only = [s["prefix_cot"] + " " + s["target"] + eos for s in real_samples] + # 即训练 label = "Reasoning: \nAction: " +else: + targets_only = ["Action: " + s["target"] + eos for s in real_samples] +``` + +**当前问题**:CoT 是一个全局开关 (`use_cot=True/False`),没有支持 **部分权重** 融合。开启 CoT 后,**全部 label tokens 的 loss 权重相同**,CoT 推理部分和动作部分共享同一个 advantage。 + +#### 推荐的 CoT Loss 加权融合方案 + +**目标**:`total_loss = (1 - cot_weight) * action_loss + cot_weight * cot_loss`,其中 `cot_weight=0.1`。 + +**方案:在 `action_mask` 层面分离 CoT tokens 和 Action tokens** + +代码修改点在 `priorzero_datafactory.py:make_llm_train_samples()`: + +```python +# 在 line 336 处(action_mask 构建后),添加 CoT/Action 分离逻辑 + +if self.use_cot and hasattr(self.args, 'cot_loss_weight') and self.args.cot_loss_weight > 0: + cot_weight = self.args.cot_loss_weight # e.g., 0.1 + + # 需要 label_ids_no_cots 信息,可在 build_llm_samples 中额外存储 + # 构建两套 mask + cot_action_mask = action_mask.clone() # 全部 label tokens + pure_action_mask = torch.zeros_like(action_mask) + + for i, (tgt_ids, tgt_no_cot_ids) in enumerate(zip(tgt_ids_list, label_ids_no_cots_list)): + no_cot_len = len(tgt_no_cot_ids) + # pure_action_mask 只标记 Action 部分的 tokens + pure_action_mask[i, -no_cot_len:] = action_mask[i, -no_cot_len:] + + # 加权 mask: CoT tokens 的 mask 值 = cot_weight, Action tokens 的 mask 值 = 1.0 + weighted_action_mask = pure_action_mask.float() + (cot_action_mask - pure_action_mask).float() * cot_weight + action_mask = weighted_action_mask +``` + +这样在 `PolicyLoss.forward()` 的 `masked_mean(loss, action_mask)` 中,CoT tokens 的 loss 自然被降权到 0.1。 + +**更优雅的方案**:在 `BatchPPOTrainer.train_batch()` 中分开计算两个 loss 再加权,但需要传递额外的 `cot_mask`,改动量更大。 + +**配置添加**(在 `priorzero_config.py` 的 `PriorZeroLLMConfig` 中): +```python +cot_loss_weight: float = 0.0 # 0 = 不使用 CoT loss, > 0 = CoT tokens 的 loss 权重 +``` + +#### 需要同步修改的位置 + +1. `priorzero_config.py`: 添加 `cot_loss_weight` 字段 +2. `priorzero_datafactory.py:make_llm_train_samples()`: 构建加权 `action_mask` +3. **注意**:需要保留 `label_ids_no_cots` 信息到 `make_llm_train_samples()` 阶段。当前代码在 `_score_labels_with_prompt_logprobs()` 中有 `l_no_cots_lens`(`datafactory.py:639`),但未传递到训练样本中。需要在 `build_llm_samples()` 中额外存储 `label_ids_no_cots`。 + +--- + +## 附录:关键变量速查表 + +| 变量/函数 | 文件:行号 | 说明 | +|-----------|----------|------| +| `AdaptiveValueNormalizer.clear()` | `stability_optimizer.py:146` | 重置所有 EMA 统计量 | +| `AdaptiveValueNormalizer.normalize()` | `stability_optimizer.py:83` | clip → batch_stats → EMA update → normalize | +| `PolicyLoss.forward()` | `loss.py:44` | PPO/GSPO loss + IS correction + clipfrac | +| `BatchPPOTrainer.__init__()` | `actor.py:221` | 初始化时传入 `enable_vllm_is_correction`, `vllm_is_truncated_threshold` | +| `BatchPPOTrainer.train_batch()` | `actor.py:251` | 微批次循环,累积梯度,含扩展指标 | +| `Actor.__init__()` | `actor.py:68` | VL 模型自动检测 (`AutoConfig` + `AutoModelForVision2Seq`) | +| `Actor.forward()` | `actor.py:135` | logits→float32→log_probs→action_log_probs | +| `PolicyModel.forward()` | `actor.py:615` | 分 chunk 推理,返回 `action_log_probs [B, T_action]` | +| `DataProcessor.make_llm_train_samples()` | `datafactory.py:272` | 返回 `(flag, data)` 元组;含去重逻辑 | +| `DataProcessor._score_labels_with_prompt_logprobs()` | `datafactory.py:606` | vLLM prompt_logprobs 提取,返回 `rollout_action_logprob` | +| `DataProcessor._build_cot_prefix_texts()` | `datafactory.py:464` | CoT reasoning prefix 生成 | +| `unique_dicts_hash()` | `datafactory.py:46` | 训练样本去重(pickle + MD5) | +| `compute_approx_kl()` | `utils.py:60` | KL 散度近似(k1/k2/k3) | +| `log_probs_from_logits()` | `utils.py:111` | logits → log_softmax(含 temperature) | +| `value_normalizer.clear()` 调用点 | `entry_sync.py:293-294` | LLM→WM 切换时清空 | +| `max_rollout_staleness` 使用点 | `entry_sync.py:280` | 控制 LLM 训练样本上限 | +| `_format_reward()` | `datafactory.py:17` | CoT 格式奖励(0/1) | +| `_normalize_vllm_weight_name()` | `actor.py:21` | vLLM 权重同步时的名称规范化 | +| `_should_skip_vllm_sync_param()` | `actor.py:28` | 跳过 LoRA adapter 参数不同步到 vLLM | diff --git a/zoo/jericho/priorzero/ensure_local_lightzero.py b/zoo/jericho/priorzero/ensure_local_lightzero.py deleted file mode 100644 index 7a697176b..000000000 --- a/zoo/jericho/priorzero/ensure_local_lightzero.py +++ /dev/null @@ -1,68 +0,0 @@ -""" -Utility module to ensure local LightZero is used across all PriorZero modules. - -This ensures PriorZero uses the local LightZero installation at: -/mnt/nfs/zhangjinouwen/puyuan/LightZero - -Usage: - Import this at the beginning of any PriorZero module: - - from ensure_local_lightzero import ensure_local_lightzero - ensure_local_lightzero() -""" - -import sys -from pathlib import Path - - -def ensure_local_lightzero(): - """ - Ensures the local LightZero path is first in sys.path. - - This allows PriorZero to use a LightZero version that has been - specifically adapted for PriorZero, rather than a globally installed version. - - Also adds the PriorZero directory to sys.path to ensure PriorZero modules - can be imported. - """ - LIGHTZERO_ROOT = Path("/mnt/nfs/zhangjinouwen/puyuan/LightZero").resolve() - PRIORZERO_DIR = Path(__file__).parent.resolve() - - if not LIGHTZERO_ROOT.exists(): - print(f"⚠️ Warning: LightZero root not found at {LIGHTZERO_ROOT}") - return False - - lightzero_str = str(LIGHTZERO_ROOT) - priorzero_str = str(PRIORZERO_DIR) - - # Remove any existing LightZero paths from sys.path - sys.path = [p for p in sys.path if 'LightZero' not in p or p == lightzero_str] - - # Insert local LightZero at the beginning - if lightzero_str not in sys.path: - sys.path.insert(0, lightzero_str) - - # Also ensure PriorZero directory is in sys.path for module imports - if priorzero_str not in sys.path: - sys.path.insert(0, priorzero_str) - - # Verify - try: - import lzero - lzero_path = Path(lzero.__file__).parent.parent - - if lzero_path == LIGHTZERO_ROOT: - print(f"✓ Using local LightZero: {lzero_path}") - print(f"✓ PriorZero modules path: {priorzero_str}") - return True - else: - print(f"⚠️ Warning: Using LightZero from {lzero_path}") - print(f" Expected: {LIGHTZERO_ROOT}") - return False - except ImportError as e: - print(f"⚠️ Warning: Could not import lzero: {e}") - return False - - -# Auto-ensure on import -ensure_local_lightzero() diff --git a/zoo/jericho/priorzero/fix_environment.sh b/zoo/jericho/priorzero/fix_environment.sh deleted file mode 100644 index 8876f54df..000000000 --- a/zoo/jericho/priorzero/fix_environment.sh +++ /dev/null @@ -1,35 +0,0 @@ -#!/bin/bash -# fix_environment.sh -# Fix numpy version conflicts and other dependency issues - -echo "==========================================" -echo "Fixing PriorZero Environment Dependencies" -echo "==========================================" - -# 1. Fix numpy version (downgrade to 1.26.4 for compatibility) -echo "" -echo "1. Fixing numpy version..." -pip install "numpy<2,>=1.24.1" --force-reinstall --no-deps - -# 2. Reinstall conflicting packages -echo "" -echo "2. Reinstalling di-engine and lightzero..." -pip install di-engine==0.5.3 --no-deps -pip install lightzero==0.2.0 --no-deps - -# 3. Verify installations -echo "" -echo "3. Verifying installations..." -python -c "import numpy; print(f'numpy version: {numpy.__version__}')" -python -c "import torch; print(f'torch version: {torch.__version__}')" -python -c "import vllm; print(f'vllm version: {vllm.__version__}')" - -echo "" -echo "==========================================" -echo "Environment fix complete!" -echo "==========================================" -echo "" -echo "Now you can run:" -echo " python priorzero_config.py" -echo " python game_segment_priorzero.py" -echo " python priorzero_entry.py --quick_test" diff --git a/zoo/jericho/priorzero/game_segment_priorzero.py b/zoo/jericho/priorzero/game_segment_priorzero.py deleted file mode 100644 index 654b93e5c..000000000 --- a/zoo/jericho/priorzero/game_segment_priorzero.py +++ /dev/null @@ -1,461 +0,0 @@ -# game_segment_priorzero.py -""" -[PRIORZERO] Enhanced Game Segment for PriorZero - -This module extends the standard GameSegment to store additional information -needed for LLM policy training (SFT + RFT). - -Key Features: -- Store MCTS policy distributions for SFT training -- Store raw text observations for LLM prompt construction -- Store LLM generated priors for analysis and debugging -- Store search values for priority calculation - -Author: PriorZero Team -Date: 2025-01-20 -""" - -import numpy as np -from typing import Optional, List, Any -from lzero.mcts.buffer.game_segment import GameSegment as OriginalGameSegment - - -class GameSegment(OriginalGameSegment): - """ - [PRIORZERO-MODIFIED] - Enhanced GameSegment that stores additional data for PriorZero training. - - New attributes: - - mcts_policy_segment: List of MCTS visit count distributions (for SFT) - - raw_obs_segment: List of raw text observations (for LLM prompts) - - llm_prior_segment: List of LLM generated text (for debugging) - - search_value_segment: List of MCTS search values (for priority) - """ - - def __init__( - self, - action_space, - game_segment_length: int = 200, - config: Optional[Any] = None, - task_id: Optional[int] = None - ): - """ - Initialize enhanced GameSegment. - - Args: - action_space: Action space from environment - game_segment_length: Maximum length of the segment - config: Policy configuration - task_id: Task ID for multi-task learning - """ - super().__init__(action_space, game_segment_length, config, task_id) - - # [PRIORZERO-NEW] Additional segments for LLM training - self.mcts_policy_segment = [] # MCTS visit count distributions - self.raw_obs_segment = [] # Raw text observations - self.llm_prior_segment = [] # LLM generated priors (for debugging) - self.search_value_segment = [] # MCTS search values - - def reset(self, init_observations: List[np.ndarray]) -> None: - """ - [PRIORZERO-MODIFIED] - Reset the segment with initial observations. - - Args: - init_observations: List of initial frame stack observations - """ - super().reset(init_observations) - - # Clear PriorZero-specific segments - self.mcts_policy_segment.clear() - self.raw_obs_segment.clear() - self.llm_prior_segment.clear() - self.search_value_segment.clear() - - def append( - self, - action: int, - obs: np.ndarray, - reward: float, - action_mask: np.ndarray, - to_play: int, - **kwargs - ) -> None: - """ - [PRIORZERO-MODIFIED] - Append a new transition to the segment. - - Args: - action: Action taken - obs: Observation received - reward: Reward received - action_mask: Valid action mask - to_play: Player ID (for multi-agent) - **kwargs: Additional arguments (timestep, chance, raw_obs_text, llm_prior_text) - """ - # [PRIORZERO-NEW] Extract PriorZero-specific kwargs before passing to parent - raw_obs_text = kwargs.pop('raw_obs_text', None) - llm_prior_text = kwargs.pop('llm_prior_text', None) - - # [DEBUG] Log first few appends to see what's being passed - if len(self.raw_obs_segment) < 3: - print(f"[SEGMENT_DEBUG] append() called: kwargs keys = {list(kwargs.keys())}") - print(f"[SEGMENT_DEBUG] raw_obs_text = {raw_obs_text[:50] if raw_obs_text else 'None'}...") - - # Call parent append with remaining kwargs - super().append(action, obs, reward, action_mask, to_play, **kwargs) - - # [PRIORZERO-NEW] Initialize placeholders for new segments - # These will be filled in by store_search_stats() - self.mcts_policy_segment.append(None) - self.search_value_segment.append(None) - - # [PRIORZERO-NEW] Store raw text observation if provided - self.raw_obs_segment.append(raw_obs_text) - - # [PRIORZERO-NEW] Store LLM prior text if provided (for debugging) - self.llm_prior_segment.append(llm_prior_text) - - def store_search_stats( - self, - root_visit_dist: List[float], - value: float, - *args, - **kwargs - ) -> None: - """ - [PRIORZERO-MODIFIED] - Store MCTS search statistics. - - This method is called after MCTS search to store the visit count - distribution and search value. These will be used for: - - SFT training: MCTS policy as supervision signal for LLM - - Priority calculation: Search value for prioritized replay - - Args: - root_visit_dist: Visit count distribution from MCTS - value: Search value from MCTS - *args: Additional positional arguments (for compatibility) - **kwargs: Additional keyword arguments (improved_policy, etc.) - """ - # [FIX] Handle NaN values - import numpy as np - if value is None or (isinstance(value, float) and np.isnan(value)): - # Use 0.0 as default for NaN values - value = 0.0 - - # Call parent method to store standard statistics - super().store_search_stats(root_visit_dist, value, *args, **kwargs) - - # [PRIORZERO-NEW] Store MCTS policy distribution - # Convert to numpy array and normalize to probability distribution - policy_array = np.array(root_visit_dist, dtype=np.float32) - - if policy_array.sum() > 0: - policy_array = policy_array / policy_array.sum() - else: - # If no visits (shouldn't happen), use uniform distribution - policy_array = np.ones_like(policy_array) / len(policy_array) - - # Update the most recent position (corresponding to last append) - if len(self.mcts_policy_segment) > 0: - self.mcts_policy_segment[-1] = policy_array - - # [PRIORZERO-NEW] Store search value - if len(self.search_value_segment) > 0: - self.search_value_segment[-1] = float(value) - - def game_segment_to_array(self) -> None: - """ - [PRIORZERO-MODIFIED] - Convert all segment lists to numpy arrays for efficient storage. - - This is called when the segment is full and ready to be stored in - the replay buffer. - """ - # Call parent method to convert standard segments - super().game_segment_to_array() - - # [PRIORZERO-NEW] Convert PriorZero-specific segments to arrays - # Use object dtype to handle variable-length arrays and None values - self.mcts_policy_segment = np.array(self.mcts_policy_segment, dtype=object) - self.search_value_segment = np.array(self.search_value_segment, dtype=np.float32) - - # For text data, keep as list (more flexible for variable-length strings) - # self.raw_obs_segment and self.llm_prior_segment remain as lists - - def get_stats(self) -> dict: - """ - [PRIORZERO-NEW] - Get statistics about this game segment. - - Returns: - stats: Dictionary of statistics - """ - stats = { - 'segment_length': len(self.reward_segment) if hasattr(self, 'reward_segment') else 0, - 'total_reward': sum(self.reward_segment) if hasattr(self, 'reward_segment') else 0, - 'num_mcts_policies': sum(1 for p in self.mcts_policy_segment if p is not None), - 'num_raw_obs': sum(1 for o in self.raw_obs_segment if o is not None), - 'num_llm_priors': sum(1 for p in self.llm_prior_segment if p is not None), - 'avg_search_value': np.mean([v for v in self.search_value_segment if v is not None]) if any(v is not None for v in self.search_value_segment) else 0.0, - } - return stats - - def get_mcts_policy_for_training(self, index: int) -> Optional[np.ndarray]: - """ - [PRIORZERO-NEW] - Get MCTS policy at a specific index for training. - - Args: - index: Index in the segment - - Returns: - policy: MCTS policy distribution, or None if not available - """ - if 0 <= index < len(self.mcts_policy_segment): - return self.mcts_policy_segment[index] - return None - - def get_raw_obs_for_training(self, index: int) -> Optional[str]: - """ - [PRIORZERO-NEW] - Get raw text observation at a specific index for training. - - Args: - index: Index in the segment - - Returns: - raw_obs: Raw text observation, or None if not available - """ - if 0 <= index < len(self.raw_obs_segment): - return self.raw_obs_segment[index] - return None - - def get_history_for_training(self, index: int, history_length: int = 5) -> List[tuple]: - """ - [PRIORZERO-NEW] - Get history context for LLM prompting. - - Args: - index: Current index in the segment - history_length: Number of past transitions to include - - Returns: - history: List of (obs, action, reward) tuples - """ - history = [] - - # Get recent transitions - start_idx = max(0, index - history_length) - for i in range(start_idx, index): - if i < len(self.raw_obs_segment) and i < len(self.action_segment) and i < len(self.reward_segment): - obs_text = self.raw_obs_segment[i] - action_id = self.action_segment[i] - reward = self.reward_segment[i] - - # Only add if observation is available - if obs_text is not None: - history.append((obs_text, action_id, reward)) - - return history - - def __repr__(self) -> str: - """ - [PRIORZERO-MODIFIED] - String representation with PriorZero statistics. - """ - base_repr = super().__repr__() - stats = self.get_stats() - - priorzero_info = ( - f"\n MCTS policies: {stats['num_mcts_policies']}" - f"\n Raw observations: {stats['num_raw_obs']}" - f"\n LLM priors: {stats['num_llm_priors']}" - f"\n Avg search value: {stats['avg_search_value']:.3f}" - ) - - return base_repr + priorzero_info - - -# ============================================================================== -# Utility Functions -# ============================================================================== - -def create_priorzero_game_segment( - action_space, - game_segment_length: int = 200, - config: Optional[Any] = None, - task_id: Optional[int] = None -) -> GameSegment: - """ - Factory function to create a PriorZero GameSegment. - - Args: - action_space: Action space from environment - game_segment_length: Maximum length of the segment - config: Policy configuration - task_id: Task ID for multi-task learning - - Returns: - segment: PriorZero GameSegment instance - """ - return GameSegment(action_space, game_segment_length, config, task_id) - - -def validate_game_segment(segment: GameSegment) -> bool: - """ - Validate that a GameSegment has consistent data. - - Args: - segment: GameSegment to validate - - Returns: - is_valid: True if segment is valid, False otherwise - """ - try: - # Check basic lengths - if not hasattr(segment, 'obs_segment'): - return False - - base_length = len(segment.obs_segment) - - # Check that all segments have compatible lengths - if hasattr(segment, 'action_segment'): - if len(segment.action_segment) != base_length: - return False - - if hasattr(segment, 'reward_segment'): - if len(segment.reward_segment) != base_length: - return False - - # Check PriorZero-specific segments - if len(segment.mcts_policy_segment) != base_length: - return False - - if len(segment.raw_obs_segment) != base_length: - return False - - # Check that MCTS policies are valid when present - for policy in segment.mcts_policy_segment: - if policy is not None: - if not isinstance(policy, np.ndarray): - return False - if policy.sum() < 0.99 or policy.sum() > 1.01: # Should sum to ~1.0 - return False - if np.any(policy < 0): # Should be non-negative - return False - - return True - - except Exception as e: - print(f"Validation error: {e}") - return False - - -# ============================================================================== -# Example Usage and Testing -# ============================================================================== - -if __name__ == "__main__": - print("="*80) - print("Testing PriorZero GameSegment") - print("="*80) - - # Create a mock action space - class MockActionSpace: - def __init__(self, n): - self.n = n - - # Create a mock config with all required attributes - class MockConfig: - def __init__(self): - self.num_unroll_steps = 10 - self.td_steps = 5 - self.discount_factor = 0.99 - self.gray_scale = False - self.transform2string = False - self.sampled_algo = False - self.gumbel_algo = False - self.use_ture_chance_label_in_chance_encoder = False - self.model = type('obj', (object,), { - 'frame_stack_num': 4, - 'action_space_size': 10, - 'observation_shape': (84, 84, 3), - 'image_channel': 3 - })() - - action_space = MockActionSpace(n=10) - mock_config = MockConfig() - - # Create a game segment - segment = GameSegment(action_space, game_segment_length=100, config=mock_config) - - # Reset with initial observations - init_obs = [np.zeros((84, 84, 3)) for _ in range(4)] - segment.reset(init_obs) - - print("\n1. Empty segment:") - print(f" Length: {len(segment.obs_segment)}") - print(f" MCTS policies: {len(segment.mcts_policy_segment)}") - - # Simulate some transitions - print("\n2. Adding transitions...") - for i in range(5): - obs = np.random.rand(84, 84, 3) - action = np.random.randint(0, 10) - reward = np.random.randn() - action_mask = np.ones(10) - - # Append transition - segment.append( - action, obs, reward, action_mask, to_play=0, - raw_obs_text=f"You see a room. Step {i}.", - llm_prior_text=f"Top actions: go north, take key" - ) - - # Store MCTS stats - visit_dist = np.random.dirichlet([1.0] * 10).tolist() - value = np.random.randn() - segment.store_search_stats(visit_dist, value) - - print(f" Added {len(segment.obs_segment)} transitions") - - # Get statistics - print("\n3. Segment statistics:") - stats = segment.get_stats() - for key, value in stats.items(): - print(f" {key}: {value}") - - # Test retrieval functions - print("\n4. Testing retrieval functions:") - mcts_policy = segment.get_mcts_policy_for_training(2) - print(f" MCTS policy at index 2: {mcts_policy is not None}") - if mcts_policy is not None: - print(f" Shape: {mcts_policy.shape}") - print(f" Sum: {mcts_policy.sum():.3f}") - - raw_obs = segment.get_raw_obs_for_training(2) - print(f" Raw obs at index 2: {raw_obs}") - - history = segment.get_history_for_training(4, history_length=3) - print(f" History for index 4: {len(history)} transitions") - - # Validate segment - print("\n5. Validating segment:") - is_valid = validate_game_segment(segment) - print(f" Is valid: {is_valid}") - - # Convert to array - print("\n6. Converting to array:") - segment.game_segment_to_array() - print(f" MCTS policy type: {type(segment.mcts_policy_segment)}") - print(f" Search value type: {type(segment.search_value_segment)}") - - # Print representation - print("\n7. Segment representation:") - print(segment) - - print("\n" + "="*80) - print("✓ All tests passed!") - print("="*80) diff --git a/zoo/jericho/priorzero/prior_generator.py b/zoo/jericho/priorzero/prior_generator.py new file mode 100644 index 000000000..3e48606c0 --- /dev/null +++ b/zoo/jericho/priorzero/prior_generator.py @@ -0,0 +1,1330 @@ +""" +Unified Prior Generator Interface + +This module provides a unified interface for generating action priors +from different types of observations (text or image). +""" +from abc import ABC, abstractmethod +from typing import List, Dict, Any, Optional, Union, Tuple +import time +import numpy as np +import torch +from PIL import Image + + +class PriorGenerator(ABC): + """ + Abstract base class for prior generators. + + Subclasses should implement generate_prior() to generate action prior + distributions from observations. + """ + + def __init__(self, model_name: str, obs_type: str): + """ + Args: + model_name: Name/path of the model + obs_type: Type of observation ('text' or 'image') + """ + self.model_name = model_name + self.obs_type = obs_type + + @abstractmethod + def generate_prior( + self, + observation: Any, + action_candidates: List[str], + history: Optional[List] = None, + temperature: float = 1.0, + **kwargs + ) -> Dict[str, Any]: + """ + Generate action prior distribution from observation. + + Args: + observation: Observation (text string or image array/PIL Image) + action_candidates: List of valid action strings + history: Optional history of previous (obs, action, reward) tuples + temperature: Temperature for sampling + **kwargs: Additional model-specific arguments + + Returns: + Dictionary containing: + - 'action_probs': np.ndarray of shape (num_actions,) with probabilities + - 'action_logits': np.ndarray of shape (num_actions,) with logits + - 'raw_output': Raw model output (for logging/debugging) + """ + pass + + @abstractmethod + def batch_generate_prior( + self, + observations: List[Any], + action_candidates_list: List[List[str]], + histories: Optional[List[List]] = None, + temperature: float = 1.0, + **kwargs + ) -> List[Dict[str, Any]]: + """ + Batch version of generate_prior for efficiency. + + Args: + observations: List of observations + action_candidates_list: List of action candidate lists + histories: Optional list of histories + temperature: Temperature for sampling + **kwargs: Additional arguments + + Returns: + List of prior dictionaries (same format as generate_prior) + """ + pass + + +class LLMPriorGenerator(PriorGenerator): + """ + Prior generator using Language Models for text observations. + + This is a wrapper around the existing vLLM engine and DataProcessor. + """ + + def __init__( + self, + vllm_engine, + data_processor, + model_name: str, + use_cot: bool = True, + **kwargs + ): + """ + Args: + vllm_engine: vLLM engine instance + data_processor: DataProcessor instance + model_name: LLM model name + use_cot: Whether to use Chain-of-Thought + """ + super().__init__(model_name, obs_type='text') + self.vllm_engine = vllm_engine + self.data_processor = data_processor + self.use_cot = use_cot + + def generate_prior( + self, + observation: str, + action_candidates: List[str], + history: Optional[List] = None, + temperature: float = 1.0, + **kwargs + ) -> Dict[str, Any]: + """ + Generate prior from text observation using LLM. + + Args: + observation: Text observation string + action_candidates: List of valid action strings + history: Optional history buffer + temperature: Sampling temperature + + Returns: + Prior dictionary with action_probs, action_logits, raw_output + """ + # Use existing DataProcessor logic + # This delegates to the existing implementation + result = self.data_processor.get_action_prior_single( + text_obs=observation, + action_candidates=action_candidates, + history=history, + temperature=temperature, + use_cot=self.use_cot, + ) + + return result + + def batch_generate_prior( + self, + observations: List[str], + action_candidates_list: List[List[str]], + histories: Optional[List[List]] = None, + temperature: float = 1.0, + **kwargs + ) -> List[Dict[str, Any]]: + """ + Batch generate priors from text observations. + """ + # Use existing DataProcessor batch logic + results = self.data_processor.get_action_prior_batch( + text_obs_list=observations, + action_candidates_list=action_candidates_list, + histories=histories, + temperature=temperature, + use_cot=self.use_cot, + ) + + return results + + +class VLPriorGenerator(PriorGenerator): + """ + Prior generator using Vision-Language (VL) models for image observations. + + Supports models like Qwen-VL, LLaVA, InternVL, etc. + + Includes training sample construction for PPO optimization with advantages. + """ + + def __init__( + self, + vl_engine, + model_name: str, + use_cot: bool = True, + tokenizer=None, + game_description: str = "", + vlm_image_mode: str = "current_only", + prompt_style: str = "concise", + logprob_extraction_mode: str = "exact", + **kwargs + ): + """ + Args: + vl_engine: VL engine instance (to be implemented) + model_name: VL model name + use_cot: Whether to use Chain-of-Thought reasoning + tokenizer: Tokenizer for building training samples + game_description: Game-specific description for prompts + vlm_image_mode: Image mode - "current_only", "first_and_current", or "all_history" + prompt_style: "concise" (shorter, better for small VLMs) or "legacy" (verbose, original) + logprob_extraction_mode: "exact" (LLM-aligned, default) or "approximate" (fallback with pseudo logprobs) + """ + super().__init__(model_name, obs_type='image') + self.vl_engine = vl_engine + self.use_cot = use_cot + self.tokenizer = tokenizer + self.game_description = game_description + self.vlm_image_mode = vlm_image_mode + self.prompt_style = prompt_style + self.logprob_extraction_mode = logprob_extraction_mode + + # For logging VL outputs + self.episode_output = [] + + # Log control: only log every N calls + self.log_interval = 100 # Log every 100 calls + self.call_count = 0 + self.batch_call_count = 0 + + def _convert_obs_to_pil_image(self, obs: np.ndarray) -> Image.Image: + """ + Robustly convert observation array to PIL Image. + + Handles various input formats: + - CHW format (C, H, W): channels first, e.g., (3, 64, 64) + - HWC format (H, W, C): channels last, e.g., (64, 64, 3) + - Grayscale (H, W): single channel, e.g., (64, 64) + - Stacked frames (N, H, W): takes the last frame + + Args: + obs: Observation array + + Returns: + PIL Image in RGB format + + Raises: + ValueError: If observation shape is invalid + """ + if not isinstance(obs, np.ndarray): + raise TypeError(f"Expected np.ndarray, got {type(obs)}") + + # Ensure uint8 dtype + if obs.dtype != np.uint8: + # Normalize to [0, 255] if needed + if obs.max() <= 1.0: + obs = (obs * 255).astype(np.uint8) + else: + obs = obs.astype(np.uint8) + + # Handle different shapes + if obs.ndim == 2: + # Grayscale (H, W) -> convert to RGB + return Image.fromarray(obs, mode='L').convert('RGB') + + elif obs.ndim == 3: + # Determine if CHW or HWC format + c, h, w = obs.shape + + # If first dimension is small (1-4), likely CHW format + if c <= 4 and h > c and w > c: + # CHW format -> transpose to HWC + if c == 1: + # Single channel (1, H, W) -> (H, W) + obs = obs[0] + return Image.fromarray(obs, mode='L').convert('RGB') + elif c == 3: + # RGB (3, H, W) -> (H, W, 3) + obs = np.transpose(obs, (1, 2, 0)) + return Image.fromarray(obs) + elif c == 4: + # RGBA or stacked frames + # Take last 3 channels as RGB + obs = np.transpose(obs[-3:], (1, 2, 0)) + return Image.fromarray(obs) + else: + # Stacked grayscale frames (N, H, W) -> take last frame + obs = obs[-1] + return Image.fromarray(obs, mode='L').convert('RGB') + + # Otherwise, assume HWC format + elif w <= 4 and h > w and c > w: + # HWC format + if w == 1: + # Single channel (H, W, 1) -> (H, W) + obs = obs[:, :, 0] + return Image.fromarray(obs, mode='L').convert('RGB') + elif w == 3: + # RGB (H, W, 3) + return Image.fromarray(obs) + elif w == 4: + # RGBA (H, W, 4) -> take first 3 channels + obs = obs[:, :, :3] + return Image.fromarray(obs) + + # Ambiguous shape - provide detailed error + raise ValueError( + f"Cannot determine image format from shape {obs.shape}. " + f"Expected CHW (C, H, W) with C<=4 or HWC (H, W, C) with C<=4. " + f"Please ensure observation is in correct format." + ) + + elif obs.ndim == 4: + # Batch dimension (B, C, H, W) or (B, H, W, C) -> take first image + raise ValueError( + f"Observation has batch dimension {obs.shape}. " + f"Please pass individual observations, not batches." + ) + + else: + raise ValueError( + f"Invalid observation shape {obs.shape}. " + f"Expected 2D (H, W) or 3D (C, H, W) or (H, W, C)." + ) + + def _assemble_images( + self, + current_obs: Union[np.ndarray, Image.Image], + history: Optional[List] = None, + ) -> List[Image.Image]: + """ + Assemble image list based on vlm_image_mode. + + Args: + current_obs: Current frame observation + history: History entries, each is (raw_obs, action, reward, timestep) + + Returns: + List of PIL Images to send to the VL model + """ + if self.vlm_image_mode == "current_only": + current_image = self._convert_obs_to_pil_image(current_obs) if isinstance(current_obs, np.ndarray) else current_obs + return [current_image] + + # Extract history images + history_images = [] + if history: + for entry in history: + obs = entry[0] # (raw_obs, action, reward, timestep) + if isinstance(obs, np.ndarray): + history_images.append(self._convert_obs_to_pil_image(obs)) + elif isinstance(obs, Image.Image): + history_images.append(obs) + # Skip non-image observations (e.g. text strings) + + current_image = self._convert_obs_to_pil_image(current_obs) if isinstance(current_obs, np.ndarray) else current_obs + + if self.vlm_image_mode == "first_and_current": + if history_images: + return [history_images[0], current_image] + return [current_image] + + elif self.vlm_image_mode == "all_history": + return history_images + [current_image] + + # Fallback (should not reach here due to validation) + return [current_image] + + def get_system_prompt(self) -> str: + """System prompt — dispatches to concise or legacy style.""" + if self.prompt_style == "concise": + return self._get_system_prompt_concise() + return self._get_system_prompt_legacy() + + def _get_system_prompt_concise(self) -> str: + """Short system prompt optimized for small VLMs (2B-7B).""" + if self.use_cot: + return ( + "You play an image-based game. Pick the best action.\n" + "Reply EXACTLY:\nReasoning: <1 sentence>\nAction: " + ) + return "You play an image-based game. Pick the best action.\nReply EXACTLY:\nAction: " + + def _get_system_prompt_legacy(self) -> str: + """ + System prompt for VL — mirrors LLM's get_system_prompt(), + only replacing "text-based adventure game" with image-based context. + """ + parts = [ + "You are an expert player in an image-based game. Your goal is to maximize the score by choosing the optimal next action.", + "Analyze the game screen and history to decide the single best next action.", + "IMPORTANT: You MUST choose EXACTLY ONE action from the provided valid actions list. Output the action name EXACTLY as given.", + ] + + if self.use_cot: + parts.append( + "OUTPUT FORMAT (you MUST follow this EXACTLY):\n" + "Reasoning: \n" + "Action: \n\n" + "RULES:\n" + "- Keep reasoning SHORT (1-3 sentences max).\n" + "- The Action line MUST contain exactly one action name from the valid actions list.\n" + "- Do NOT add any text after the action name." + ) + else: + parts.append( + "OUTPUT FORMAT:\n" + "Action: \n\n" + "Output ONLY this single line. No other text." + ) + return "\n".join(parts) + + def get_user_prompt( + self, + action_candidates: List[str], + history: Optional[List] = None, + num_images: int = 1, + ) -> str: + """User prompt — dispatches to concise or legacy style.""" + if self.prompt_style == "concise": + return self._get_user_prompt_concise(action_candidates, history, num_images) + return self._get_user_prompt_legacy(action_candidates, history, num_images) + + def _get_user_prompt_concise( + self, + action_candidates: List[str], + history: Optional[List] = None, + num_images: int = 1, + ) -> str: + """ + Concise user prompt: minimal tokens, maximum signal. + Designed for small VLMs (2B-7B) where instruction-following degrades with long prompts. + """ + parts = [] + + # Game description — one line only + if self.game_description: + # Take only the first sentence of game_description + first_sentence = self.game_description.split('\n')[0].strip() + parts.append(first_sentence) + + # Multi-image labelling + if self.vlm_image_mode != "current_only" and num_images > 1: + img_idx = 1 + if history and len(history) > 0: + for entry in history: + action = entry[1] + reward = entry[2] + has_image = isinstance(entry[0], (np.ndarray, Image.Image)) + if self.vlm_image_mode == "all_history" and has_image and img_idx < num_images: + parts.append(f"[Image {img_idx}] Action: {action}, Reward: {reward}") + img_idx += 1 + elif self.vlm_image_mode == "first_and_current" and has_image and img_idx == 1: + parts.append(f"[Image {img_idx} - initial] Action: {action}, Reward: {reward}") + img_idx += 1 + else: + parts.append(f"Action: {action}, R: {reward}") + parts.append(f"[Image {num_images}] Current screen.") + else: + # Single image — text-only history + if history and len(history) > 0: + hist_strs = [] + for entry in history: + action, reward = entry[1], entry[2] + hist_strs.append(f"{action}(R:{reward})") + parts.append("History: " + " → ".join(hist_strs)) + parts.append("Current screen shown above.") + + # Valid actions — compact + actions_str = ", ".join(action_candidates) + parts.append(f"Actions: [{actions_str}]") + + # LunarLander-specific compact hints + if set(action_candidates) == {"NOOP", "LEFT_ENGINE", "MAIN_ENGINE", "RIGHT_ENGINE"}: + parts.append( + "NOOP=do nothing | LEFT_ENGINE=push right,rotate CW(-0.03) | " + "MAIN_ENGINE=slow descent(-0.3) | RIGHT_ENGINE=push left,rotate CCW(-0.03)\n" + "Goal: land on pad horizontally. Crash=-100, land=+100." + ) + + # Instruction + if self.use_cot: + parts.append("Reasoning: <1 sentence>\nAction: ") + else: + parts.append("Action: ") + + return "\n".join(parts) + + def _get_user_prompt_legacy( + self, + action_candidates: List[str], + history: Optional[List] = None, + num_images: int = 1, + ) -> str: + """ + User prompt for VL — mirrors LLM's get_user_prompt() structure, + replacing text observation with image vision tokens. + + Args: + action_candidates: List of valid action names + history: Optional history entries + num_images: Number of images being sent (for multi-image labelling) + """ + prompt_parts = [] + + # Multi-image mode: label each image in the prompt + if self.vlm_image_mode != "current_only" and num_images > 1: + img_idx = 1 # 1-based image index for the prompt + + if history and len(history) > 0: + prompt_parts.append("=== GAME HISTORY ===") + for entry in history: + if len(entry) >= 4: + obs, action, reward, timestep = entry[0], entry[1], entry[2], entry[3] + else: + obs, action, reward = entry[0], entry[1], entry[2] + timestep = None + + # Check if this history entry has a corresponding image + has_image = isinstance(obs, (np.ndarray, Image.Image)) + + if self.vlm_image_mode == "all_history" and has_image and img_idx < num_images: + step_label = f"Step {timestep}" if timestep is not None else "Step" + prompt_parts.append(f"=== HISTORICAL OBSERVATION ({step_label}) ===") + prompt_parts.append(f"[See image {img_idx} above]") + if timestep is not None: + prompt_parts.append(f"Action: {action}, Reward: {reward}") + else: + prompt_parts.append(f"Action: {action}, Reward: {reward}") + img_idx += 1 + elif self.vlm_image_mode == "first_and_current" and has_image and img_idx == 1: + step_label = f"Step {timestep}" if timestep is not None else "First Step" + prompt_parts.append(f"=== INITIAL OBSERVATION ({step_label}) ===") + prompt_parts.append(f"[See image {img_idx} above]") + prompt_parts.append(f"Action: {action}, Reward: {reward}") + img_idx += 1 + else: + # Text-only history entry + if timestep is not None: + prompt_parts.append(f"Step {timestep}: Action: {action}, Reward: {reward}") + else: + prompt_parts.append(f"Action: {action}, Reward: {reward}") + + prompt_parts.append("") # empty line separator + + prompt_parts.append("=== CURRENT OBSERVATION ===") + prompt_parts.append(f"[See image {num_images} above]") + prompt_parts.append("\nLook at the image carefully and analyze:") + prompt_parts.append("- The lander's tilt angle (horizontal, tilted left, or tilted right?)") + prompt_parts.append("- The lander's horizontal position relative to the landing pad") + prompt_parts.append("- Visual indicators of descent speed") + + else: + # Original single-image prompt (current_only mode or only 1 image) + if history and len(history) > 0: + prompt_parts.append("=== GAME HISTORY ===") + for entry in history: + if len(entry) >= 4: + obs, action, reward, timestep = entry[0], entry[1], entry[2], entry[3] + prompt_parts.append(f"Step {timestep}: Action: {action}, Reward: {reward}") + else: + obs, action, reward = entry[0], entry[1], entry[2] + prompt_parts.append(f"Action: {action}, Reward: {reward}") + prompt_parts.append("") # empty line separator + + prompt_parts.append("=== CURRENT OBSERVATION ===") + prompt_parts.append("[See the game screen image above]") + prompt_parts.append("\nLook at the image carefully and analyze:") + prompt_parts.append("- The lander's tilt angle (horizontal, tilted left, or tilted right?)") + prompt_parts.append("- The lander's horizontal position relative to the landing pad") + prompt_parts.append("- Visual indicators of descent speed") + if self.game_description: + prompt_parts.append(self.game_description) + + if action_candidates and len(action_candidates) > 0: + actions_str = ", ".join(action_candidates) + prompt_parts.append(f"\nValid actions: [{actions_str}]") + + # Add per-action descriptions for LunarLander + # (For other games, the game_description already covers action semantics) + if set(action_candidates) == {"NOOP", "LEFT_ENGINE", "MAIN_ENGINE", "RIGHT_ENGINE"}: + prompt_parts.append( + "- NOOP: Do nothing (0 cost).\n" + "- LEFT_ENGINE: Fires the left thruster. Pushes the lander RIGHT and rotates it clockwise. (-0.03 cost)\n" + "- MAIN_ENGINE: Fires the bottom thruster. Slows descent. (-0.3 cost)\n" + "- RIGHT_ENGINE: Fires the right thruster. Pushes the lander LEFT and rotates it counter-clockwise. (-0.03 cost)\n" + "\n" + "=== STRATEGY GUIDE ===\n" + "1. Keep Horizontal: The game penalizes tilt. Correct tilt immediately. If tilted left, fire LEFT_ENGINE to rotate clockwise. If tilted right, fire RIGHT_ENGINE.\n" + "2. Conserve Main Fuel: MAIN_ENGINE is very expensive (-0.3). Use it ONLY if falling too fast.\n" + "3. Steer to Center: Use side engines to adjust horizontal position toward the flags.\n" + "4. Coasting: If the lander is horizontal, aligned with the pad, and descending slowly, use NOOP to save points." + ) + + prompt_parts.append("\n=== INSTRUCTION ===") + if self.use_cot: + prompt_parts.append( + "Choose the best action. You MUST respond in EXACTLY this format:\n" + "Reasoning: \n" + "Action: \n" + "\n" + "CRITICAL RULES:\n" + "- Write ONLY the action name after 'Action:', nothing else.\n" + "- Do NOT add punctuation, arrows (->), or explanations after the action.\n" + "- Do NOT write 'NOPE' or 'NO_OP', only 'NOOP'.\n" + "\n" + "Example 1:\n" + "Reasoning: The lander is tilted left and drifting left of the pad; firing LEFT_ENGINE will rotate it clockwise back to horizontal and push it right toward the center at a low cost.\n" + "Action: LEFT_ENGINE\n" + "\n" + "Example 2:\n" + "Reasoning: The lander is horizontal and centered, but falling too rapidly; despite the high cost, MAIN_ENGINE is strictly necessary to slow the descent and prevent a -100 crash penalty.\n" + "Action: MAIN_ENGINE\n" + "\n" + "Example 3:\n" + "Reasoning: The lander is perfectly horizontal, aligned above the pad, and descending at a safe, slow speed; no thrust is needed, so doing nothing avoids point deductions.\n" + "Action: NOOP" + ) + else: + example_action = action_candidates[1] if len(action_candidates) >= 2 else (action_candidates[0] if action_candidates else "NOOP") + prompt_parts.append( + f"Choose the best action. Output ONLY:\n" + f"Action: \n\n" + f"Example:\nAction: {example_action}" + ) + return "\n".join(prompt_parts) + + def _parse_vl_output_with_cot( + self, + raw_output: str, + action_candidates: List[str] + ) -> Tuple[str, Optional[str]]: + """ + Parse VL output to extract action and optional CoT reasoning. + + Args: + raw_output: Raw VL output string + action_candidates: List of valid action names + + Returns: + Tuple of (chosen_action, cot_prefix) + - chosen_action: The selected action name + - cot_prefix: The reasoning part (if use_cot=True), else None + """ + import re + + cot_prefix = None + chosen_action = None + + # Extract reasoning part (if present) + if self.use_cot: + reasoning_match = re.search(r'Reasoning:\s*(.+?)(?=Action:|$)', raw_output, re.DOTALL | re.IGNORECASE) + if reasoning_match: + cot_prefix = reasoning_match.group(1).strip() + + # Strategy 1: Extract text after "Action:" and match against candidates + # Use .+ instead of \S+ to capture multi-word or underscore-separated actions + action_match = re.search(r'Action:\s*(.+)', raw_output, re.IGNORECASE) + if action_match: + action_str = action_match.group(1).strip().strip("'\"`.,:;") + # Exact match (case-insensitive) + for candidate in action_candidates: + if candidate.upper() == action_str.upper(): + chosen_action = candidate + break + # If no exact match, try if candidate is contained in the extracted text + if chosen_action is None: + for candidate in action_candidates: + if candidate.upper() in action_str.upper(): + chosen_action = candidate + break + + # Strategy 2: If no "Action:" line found, scan entire output for action names + if chosen_action is None: + # Search for exact action name mentions in the output (prefer later mentions) + last_found = None + for candidate in action_candidates: + # Use word boundary to avoid partial matches + pattern = re.escape(candidate) + matches = list(re.finditer(pattern, raw_output, re.IGNORECASE)) + if matches: + pos = matches[-1].start() + if last_found is None or pos > last_found[1]: + last_found = (candidate, pos) + if last_found is not None: + chosen_action = last_found[0] + + # Fallback: if no valid action found, use first candidate + if chosen_action is None: + chosen_action = action_candidates[0] if action_candidates else "NOOP" + + return chosen_action, cot_prefix + + def _extract_action_logprobs_batch( + self, + image_list: List[Image.Image], + prompt: str, + action_candidates: List[str], + cot_prefix: Optional[str], + temperature: float = 1.0 + ) -> Tuple[Optional[np.ndarray], Dict[str, List], Dict[str, List], Dict[str, List]]: + """ + Extract action log probabilities with configurable mode. + """ + if self.logprob_extraction_mode == "exact": + return self._extract_logprobs_exact_mode( + image_list, prompt, action_candidates, cot_prefix, temperature + ) + else: # approximate mode (default) + return self._extract_logprobs_approximate_mode( + image_list, prompt, action_candidates, cot_prefix, temperature + ) + + def _extract_logprobs_approximate_mode( + self, + image_list: List[Image.Image], + prompt: str, + action_candidates: List[str], + cot_prefix: Optional[str], + temperature: float = 1.0 + ) -> Tuple[Optional[np.ndarray], Dict[str, List], Dict[str, List], Dict[str, List]]: + """ + Approximate mode: Use fallback with pseudo token data. + Fast but less accurate. + """ + import logging + logger = logging.getLogger(__name__) + + try: + from transformers import AutoTokenizer + tokenizer = AutoTokenizer.from_pretrained(self.model_name, trust_remote_code=True) + + rollout_action_logprob_dict = {} + full_ids_dict = {} + label_ids_dict = {} + + for action in action_candidates: + if self.use_cot and cot_prefix: + label_text = cot_prefix + " " + action + else: + label_text = "Action: " + action + + label_ids = tokenizer(label_text, add_special_tokens=False)["input_ids"] + full_prompt = prompt + "\n" + label_text + full_ids = tokenizer(full_prompt, add_special_tokens=False)["input_ids"] + + pseudo_logprobs = [0.0] * len(label_ids) + + rollout_action_logprob_dict[action] = pseudo_logprobs + full_ids_dict[action] = full_ids + label_ids_dict[action] = label_ids + + return None, rollout_action_logprob_dict, full_ids_dict, label_ids_dict + + except Exception as e: + logger.error(f"⚠️ Approximate mode failed: {e}", exc_info=True) + + return None, {}, {}, {} + + def _extract_logprobs_exact_mode( + self, + image_list: List[Image.Image], + prompt: str, + action_candidates: List[str], + cot_prefix: Optional[str], + temperature: float = 1.0 + ) -> Tuple[Optional[np.ndarray], Dict[str, List], Dict[str, List], Dict[str, List]]: + """ + Exact mode: Use token IDs like LLM (bypassing chat template). + """ + import logging + import math + logger = logging.getLogger(__name__) + + try: + from transformers import AutoTokenizer + tokenizer = AutoTokenizer.from_pretrained(self.model_name, trust_remote_code=True) + + prompt_ids = tokenizer(prompt, add_special_tokens=False)["input_ids"] + + if self.use_cot and cot_prefix: + label_texts = [cot_prefix + " " + action for action in action_candidates] + label_texts_no_cots = [" " + action for action in action_candidates] + else: + label_texts = ["Action: " + action for action in action_candidates] + label_texts_no_cots = label_texts + + label_ids_list = [tokenizer(label, add_special_tokens=False)["input_ids"] for label in label_texts] + label_ids_no_cots_list = [tokenizer(label, add_special_tokens=False)["input_ids"] for label in label_texts_no_cots] + full_ids_list = [prompt_ids + label_ids for label_ids in label_ids_list] + + results = self.vl_engine.batch_generate_with_token_ids( + images=[image_list] * len(action_candidates), + prompt_token_ids=full_ids_list, + temperature=temperature, + max_new_tokens=1, + return_logprobs=True, + ) + + action_scores = [] + rollout_action_logprob_dict = {} + full_ids_dict = {} + label_ids_dict = {} + + for action, label_ids, label_ids_no_cot, full_ids, result in zip( + action_candidates, label_ids_list, label_ids_no_cots_list, full_ids_list, results + ): + prompt_logprobs = result.get('prompt_logprobs') if isinstance(result, dict) else None + + if not prompt_logprobs or len(prompt_logprobs) == 0: + action_scores.append(float("-inf")) + rollout_action_logprob_dict[action] = [] + full_ids_dict[action] = [] + label_ids_dict[action] = [] + continue + + token_lps = [] + for j in range(1, len(full_ids)): + tok_id = full_ids[j] + lp_dict = prompt_logprobs[j] + + if lp_dict is None or tok_id not in lp_dict: + break + + logprob_obj = lp_dict[tok_id] + logprob = logprob_obj.logprob if hasattr(logprob_obj, 'logprob') else float(logprob_obj) + + if math.isnan(logprob): + break + + token_lps.append(logprob) + + if len(token_lps) > 0: + l_len = len(label_ids) + l_no_cots_len = len(label_ids_no_cot) + label_lps = token_lps[-l_len:] + + if self.use_cot: + target_lps = label_lps + else: + target_lps = label_lps[-l_no_cots_len:] + + score = sum(target_lps) / len(target_lps) + action_scores.append(score) + rollout_action_logprob_dict[action] = label_lps + full_ids_dict[action] = full_ids + label_ids_dict[action] = label_ids + else: + action_scores.append(float("-inf")) + rollout_action_logprob_dict[action] = [] + full_ids_dict[action] = [] + label_ids_dict[action] = [] + + valid_count = sum(1 for s in action_scores if s > float("-inf")) + if valid_count == len(action_candidates): + return np.array(action_scores, dtype=np.float32), rollout_action_logprob_dict, full_ids_dict, label_ids_dict + + logger.warning(f"⚠️ Exact mode: {valid_count}/{len(action_candidates)} valid") + + except Exception as e: + logger.error(f"⚠️ Exact mode failed: {e}") + + return None, {}, {}, {} + + def _action_to_logprob( + self, + chosen_action: str, + action_candidates: List[str], + temperature: float = 1.0 + ) -> np.ndarray: + """ + Convert chosen action to log probability distribution. + + For training, we need to store the "old" log probabilities that were used + to select the action. This creates a peaked distribution around the chosen action. + + Args: + chosen_action: The action selected by VL + action_candidates: List of all valid actions + temperature: Temperature for softening the distribution + + Returns: + Log probability array of shape (num_actions,) + """ + num_actions = len(action_candidates) + + # Create peaked but NOT one-hot distribution to preserve MCTS exploration. + # Use moderate logit gap (2.0 vs 0.0) instead of extreme (10.0 vs -10.0), + # so the prior is informative but not deterministic. + logits = np.zeros(num_actions, dtype=np.float32) + + try: + chosen_idx = action_candidates.index(chosen_action) + logits[chosen_idx] = 2.0 # Moderate logit for chosen action + except ValueError: + # If chosen action not in candidates, uniform distribution + logits = np.zeros(num_actions) + + # Apply temperature and convert to log probabilities (numerically stable) + logits = logits / temperature + max_logit = np.max(logits) + log_probs = logits - max_logit - np.log(np.sum(np.exp(logits - max_logit)) + 1e-10) + + return log_probs + + def _parse_vl_output( + self, + raw_output: str, + action_candidates: List[str] + ) -> np.ndarray: + """ + Parse VL output to extract action probabilities. + + Args: + raw_output: Raw text output from VL + action_candidates: List of valid action names (e.g., ['NOOP', 'FIRE', 'RIGHT']) + + Returns: + Action probabilities as numpy array + """ + import json + import re + + # Try to extract JSON from output + try: + # Look for JSON-like structure + json_match = re.search(r'\{[^}]+\}', raw_output) + if json_match: + action_probs_dict = json.loads(json_match.group()) + + # Convert to array aligned with action_candidates + probs = [] + for action in action_candidates: + # Try exact match and case-insensitive match + prob = action_probs_dict.get(action, + action_probs_dict.get(action.upper(), + action_probs_dict.get(action.lower(), 0.0))) + probs.append(prob) + + probs = np.array(probs, dtype=np.float32) + + # Normalize + if probs.sum() > 0: + probs = probs / probs.sum() + else: + # Fallback to uniform + probs = np.ones(len(action_candidates), dtype=np.float32) / len(action_candidates) + + return probs + except Exception as e: + import logging + logger = logging.getLogger(__name__) + logger.warning(f"Failed to parse VL output: {e}. Using uniform prior.") + logger.debug(f"Raw output: {raw_output}") + + # Fallback: uniform distribution + return np.ones(len(action_candidates), dtype=np.float32) / len(action_candidates) + + def generate_prior( + self, + observation: Union[np.ndarray, Image.Image], + action_candidates: List[str], + history: Optional[List] = None, + temperature: float = 1.0, + **kwargs + ) -> Dict[str, Any]: + """ + Generate prior with LLM-aligned token-level data. + """ + self.call_count += 1 + + # Assemble images + image_list = self._assemble_images(observation, history) + prompt = self.get_user_prompt(action_candidates, history, num_images=len(image_list)) + + # Step 1: Generate to get chosen action and CoT prefix + result = self.vl_engine.generate( + image=image_list, + prompt=prompt, + temperature=temperature, + system_prompt=self.get_system_prompt(), + return_logprobs=False, + **kwargs + ) + raw_output = result.get('text', '') if isinstance(result, dict) else result + chosen_action, cot_prefix = self._parse_vl_output_with_cot(raw_output, action_candidates) + + # Step 2: Extract logprobs with token-level data (same as LLM) + action_log_probs, rollout_logprob_dict, full_ids_dict, label_ids_dict = self._extract_action_logprobs_batch( + image_list, prompt, action_candidates, cot_prefix, temperature + ) + + # Fallback if batch extraction failed + if action_log_probs is None: + action_log_probs = self._action_to_logprob(chosen_action, action_candidates, temperature) + rollout_logprob_dict = {} + full_ids_dict = {} + label_ids_dict = {} + + action_probs = np.exp(action_log_probs) + + return { + 'action_probs': action_probs, + 'action_logits': action_log_probs, + 'raw_output': raw_output, + 'cot_prefix': cot_prefix, + 'chosen_action': chosen_action, + 'rollout_action_logprob': rollout_logprob_dict, + 'full_ids': full_ids_dict, + 'label_ids': label_ids_dict, + } + + def batch_generate_prior( + self, + observations: List[Union[np.ndarray, Image.Image]], + action_candidates_list: List[List[str]], + histories: Optional[List[List]] = None, + temperature: float = 1.0, + **kwargs + ) -> List[Dict[str, Any]]: + """ + Batch generate priors with LLM-aligned token-level data. + """ + if histories is None: + histories = [None] * len(observations) + + # Assemble images and prompts + image_lists = [] + prompts = [] + for obs, history, action_candidates in zip(observations, histories, action_candidates_list): + image_list = self._assemble_images(obs, history) + prompt = self.get_user_prompt(action_candidates, history, num_images=len(image_list)) + image_lists.append(image_list) + prompts.append(prompt) + + # Step 1: Generate to get chosen actions and CoT prefixes + raw_outputs = self.vl_engine.batch_generate( + images=image_lists, + prompts=prompts, + temperature=temperature, + system_prompt=self.get_system_prompt(), + return_logprobs=False, + **kwargs + ) + + # Parse outputs + chosen_actions = [] + cot_prefixes = [] + for result, action_candidates in zip(raw_outputs, action_candidates_list): + raw_output = result.get('text', '') if isinstance(result, dict) else result + chosen_action, cot_prefix = self._parse_vl_output_with_cot(raw_output, action_candidates) + chosen_actions.append(chosen_action) + cot_prefixes.append(cot_prefix) + + # Step 2: Extract logprobs with token-level data for each observation + results = [] + for idx, (image_list, prompt, action_candidates, raw_output, chosen_action, cot_prefix) in enumerate( + zip(image_lists, prompts, action_candidates_list, + [r.get('text', '') if isinstance(r, dict) else r for r in raw_outputs], + chosen_actions, cot_prefixes) + ): + action_log_probs, rollout_logprob_dict, full_ids_dict, label_ids_dict = self._extract_action_logprobs_batch( + image_list, prompt, action_candidates, cot_prefix, temperature + ) + + if action_log_probs is None: + action_log_probs = self._action_to_logprob(chosen_action, action_candidates, temperature) + rollout_logprob_dict = {} + full_ids_dict = {} + label_ids_dict = {} + + action_probs = np.exp(action_log_probs) + + results.append({ + 'action_probs': action_probs, + 'action_logits': action_log_probs, + 'raw_output': raw_output, + 'cot_prefix': cot_prefix, + 'chosen_action': chosen_action, + 'rollout_action_logprob': rollout_logprob_dict, + 'full_ids': full_ids_dict, + 'label_ids': label_ids_dict, + }) + + return results + + def build_vl_train_samples( + self, + raw_obs_list: List[List[np.ndarray]], + history_obs_list: List[List[List]], + vl_prior_per_tok_list: List[List[Dict]], + pred_values: Optional[torch.Tensor] = None, + target_values: Optional[torch.Tensor] = None, + cot_prefix_list: Optional[List[List[str]]] = None, + vl_action_list: Optional[List[List[str]]] = None, + ) -> List[Dict[str, Any]]: + """ + Build training samples for VL - ALIGNED with LLM's build_llm_samples. + + Args: + raw_obs_list: [B, T] Raw image observations + history_obs_list: [B, T] History observations + vl_prior_per_tok_list: [B, T] VL prior per token (with rollout_logprob, full_ids, label_ids) + pred_values: [B, T-1] Predicted values + target_values: [B, T-1] Target values + cot_prefix_list: [B, T] CoT prefixes + vl_action_list: [B, T] Action names + + Returns: + List of training samples + """ + import logging + logger = logging.getLogger(__name__) + + samples = [] + B = len(raw_obs_list) + if B == 0: + return samples + T = len(raw_obs_list[0]) + + for b in range(B): + for t in range(T - 1): + current_obs = raw_obs_list[b][t] + current_hist = history_obs_list[b][t] + + # Convert obs to PIL Image + if isinstance(current_obs, np.ndarray): + image = self._convert_obs_to_pil_image(current_obs) + else: + image = current_obs + + # Build prompt (same structure as LLM) + image_list = self._assemble_images(current_obs, current_hist) + instruction = self.get_user_prompt( + action_candidates=None, # Will be filled from vl_prior_per_tok_list + history=current_hist, + num_images=len(image_list) + ) + + # Get action and logprobs (same as LLM) + true_action = vl_action_list[b][t+1] + rollout_logprob = vl_prior_per_tok_list[b][t+1]['rollout_action_logprob'][true_action] + full_ids = vl_prior_per_tok_list[b][t+1]['full_ids'][true_action] + label_ids = vl_prior_per_tok_list[b][t+1]['label_ids'][true_action] + + if len(label_ids) == 0: + continue + + # Get values (same as LLM) + target_value = None + if target_values is not None: + target_value = float(target_values[b][t].item()) + + pred_value = None + if pred_values is not None: + pred_value = float(pred_values[b][t].item()) + + # Get CoT prefix (same as LLM) + prefix_cot = None + if self.use_cot and cot_prefix_list is not None: + prefix_cot = cot_prefix_list[b][t+1] + + samples.append({ + "image": image, + "image_list": image_list, + "instruction": instruction, + "target": true_action, + "pred_value": pred_value, + "target_value": target_value, + "rollout_logprob": rollout_logprob, + "prefix_cot": prefix_cot, + "full_ids": full_ids, + "label_ids": label_ids, + }) + + return samples + + def compute_action_log_prob( + self, + vl_output: str, + target_action: str, + valid_actions: List[str], + temperature: float = 1.0 + ) -> float: + """ + Compute log probability of target action from VL output. + + This is used during training to compute the new log probability + for PPO ratio calculation. + + Args: + vl_output: Raw VL output string + target_action: The action that was actually taken + valid_actions: List of valid action names + temperature: Temperature for scaling + + Returns: + Log probability of target action + """ + if self.use_cot: + # Parse CoT output to get chosen action + chosen_action, _ = self._parse_vl_output_with_cot(vl_output, valid_actions) + + # Get log prob distribution + log_probs = self._action_to_logprob(chosen_action, valid_actions, temperature) + + # Return log prob of target action + try: + target_idx = valid_actions.index(target_action) + return float(log_probs[target_idx]) + except ValueError: + # Target action not in valid actions + return -10.0 # Very low log prob + else: + # Parse probability distribution + probs = self._parse_vl_output(vl_output, valid_actions) + log_probs = np.log(probs + 1e-10) + + try: + target_idx = valid_actions.index(target_action) + return float(log_probs[target_idx]) + except ValueError: + return -10.0 + + + def get_vl_output_log( + self, + wm_train_iter: int, + vl_train_iter: int, + ) -> None: + """ + Log VL output statistics (similar to LLM's get_llm_output_log). + + Args: + wm_train_iter: World model training iteration + vl_train_iter: VL training iteration + """ + import logging + logger = logging.getLogger(__name__) + + if len(self.episode_output) == 0: + return + + logger.info( + f"\n{'='*80}\n" + f"[VL Output Log] WM Iter: {wm_train_iter} | VL Iter: {vl_train_iter}\n" + f"{'='*80}" + ) + + for i, tmp_dict in enumerate(self.episode_output[:15]): + instruction = tmp_dict["Instruction"] + response = tmp_dict["Response"] + vl_prior = tmp_dict["vl_prior_per_seq"] + chosen_action = tmp_dict.get("chosen_action", "N/A") + cot_prefix = tmp_dict.get("cot_prefix", "") + + logger.info( + f"\n{'-'*80}\n" + f"[Step {i}]\n" + f"{'-'*80}\n" + f"Instruction:\n{instruction}\n\n" + f"Response:\n{response}\n\n" + f"Chosen Action: {chosen_action}\n" + ) + + if cot_prefix: + logger.info(f"CoT Reasoning:\n{cot_prefix}\n") + + logger.info("Action Probabilities:") + + # Sort actions by probability (descending) + sorted_actions = sorted(vl_prior.items(), key=lambda x: x[1], reverse=True) + + for action, prob in sorted_actions: + logger.info(f" {action:30s} | prob={prob:.6f}") + + self.episode_output = [] + + +def create_prior_generator( + obs_type: str, + model_config: Dict[str, Any], + **kwargs +) -> PriorGenerator: + """ + Factory function to create appropriate prior generator. + + Args: + obs_type: 'text' or 'image' + model_config: Model configuration dictionary + **kwargs: Additional arguments + + Returns: + PriorGenerator instance (LLMPriorGenerator or VLPriorGenerator) + """ + if obs_type == 'text': + # Create LLM prior generator + from vllm_utils.vllm_engine import create_vllm_engine + + vllm_engine = create_vllm_engine( + tensor_parallel_size=model_config.get('tensor_parallel_size', 1), + pretrain=model_config['model_path'], + enable_prefix_caching=model_config.get('enable_prefix_caching', True), + max_model_len=model_config.get('max_model_len', 8192), + gpu_memory_utilization=model_config.get('gpu_memory_utilization', 0.3), + ) + + # Note: data_processor needs to be passed separately + # This is a placeholder - actual implementation needs data_processor + raise NotImplementedError( + "LLMPriorGenerator requires data_processor. " + "Use the existing implementation or pass data_processor explicitly." + ) + + elif obs_type == 'image': + # Create VL prior generator + from vl_engine import create_vl_engine + + vl_engine = create_vl_engine( + model_name=model_config['model_name'], + model_path=model_config['model_path'], + tensor_parallel_size=model_config.get('tensor_parallel_size', 1), + gpu_memory_utilization=model_config.get('gpu_memory_utilization', 0.3), + ) + + return VLPriorGenerator( + vl_engine=vl_engine, + model_name=model_config['model_name'], + ) + + else: + raise ValueError(f"Unknown obs_type: {obs_type}. Must be 'text' or 'image'.") + + +if __name__ == "__main__": + # Example usage + print("Prior Generator Interface") + print("=" * 80) + print("\nThis module provides unified interface for generating action priors.") + print("\nSupported generators:") + print(" - LLMPriorGenerator: For text observations (Jericho games)") + print(" - VLPriorGenerator: For image observations (Atari games)") + print("\nUsage:") + print(" generator = create_prior_generator(obs_type='image', model_config={...})") + print(" prior = generator.generate_prior(observation, action_candidates)") diff --git a/zoo/jericho/priorzero/priorzero_collector.py b/zoo/jericho/priorzero/priorzero_collector.py deleted file mode 100644 index 1fb6e53c7..000000000 --- a/zoo/jericho/priorzero/priorzero_collector.py +++ /dev/null @@ -1,800 +0,0 @@ -# priorzero_collector.py -""" -[PRIORZERO] PriorZero Collector Implementation - -This module implements async data collection with LLM prior integration. - -Key Features: -- Async LLM inference using vLLM for efficient batch generation -- History buffer management for context-aware prompting -- Error handling and retry logic for robust LLM calls -- Full alignment with UniZero collector architecture - -Author: PriorZero Team -Date: 2025-01-20 -""" - -import asyncio -import logging -import sys -import time -from collections import deque, defaultdict -from pathlib import Path -from typing import Optional, Any, List, Dict, Tuple - -# [CRITICAL] Ensure local LightZero is used -from ensure_local_lightzero import ensure_local_lightzero -ensure_local_lightzero() - -import numpy as np -import torch -from ding.envs import BaseEnvManager -from ding.torch_utils import to_ndarray -from ding.utils import build_logger, EasyTimer, SERIAL_COLLECTOR_REGISTRY -from vllm import AsyncLLMEngine, SamplingParams - -# Import from local LightZero -from lzero.worker.muzero_segment_collector import MuZeroSegmentCollector as OriginalCollector -from lzero.mcts.utils import prepare_observation -from game_segment_priorzero import GameSegment - - -# ============================================================================== -# Helper Functions -# ============================================================================== - -def extract_raw_obs_text(obs_dict: Dict[str, Any]) -> str: - """ - Extract text observation from environment observation dictionary. - - Args: - obs_dict: Observation dictionary from environment - - Returns: - text_obs: Text observation string - """ - # [PRIORZERO-FIX] Try to get 'raw_obs_text' field first (Jericho env adds this) - if 'raw_obs_text' in obs_dict: - return str(obs_dict['raw_obs_text']) - - # Try to get 'raw_obs' field (alternative naming) - if 'raw_obs' in obs_dict: - return str(obs_dict['raw_obs']) - - # Try to get 'text' field - if 'text' in obs_dict: - return str(obs_dict['text']) - - # Try to get 'observation_str' field (Jericho env provides this in save_replay mode) - if 'observation_str' in obs_dict: - return str(obs_dict['observation_str']) - - # Try to get 'observation' and check if it's text - if 'observation' in obs_dict: - obs = obs_dict['observation'] - if isinstance(obs, str): - return obs - elif isinstance(obs, (list, np.ndarray)): - # If observation is already processed (e.g., embeddings), cannot extract text - # Return a placeholder - return f"[Observation vector of shape {np.array(obs).shape}]" - - # Fallback: return str representation - return str(obs_dict) - - -# ============================================================================== -# PriorZero Collector Class -# ============================================================================== - -@SERIAL_COLLECTOR_REGISTRY.register('priorzero_segment', force_overwrite=True) -class PriorZeroCollector(OriginalCollector): - """ - [PRIORZERO-MODIFIED] - Async collector that integrates LLM priors into MCTS-based data collection. - - Features: - - Async LLM inference with vLLM engine - - History buffer for each environment (sliding window) - - Robust error handling with retries - - Detailed logging of LLM prior statistics - """ - - def __init__( - self, - vllm_engine: AsyncLLMEngine, - policy_config: Dict, - **kwargs - ): - """ - Initialize PriorZeroCollector. - - Args: - vllm_engine: vLLM async engine for LLM inference - policy_config: Policy configuration (contains llm_policy_cfg) - **kwargs: Additional arguments for parent class - """ - # [FIX] Set policy_config in kwargs before calling super().__init__ - # because parent class needs it - kwargs['policy_config'] = policy_config - - # Extract debug_mode before passing to parent (parent doesn't accept this parameter) - self.debug_mode = kwargs.pop('debug_mode', False) - - super().__init__(**kwargs) - - self.vllm_engine = vllm_engine - # self.policy_config already set by parent class from kwargs - self.llm_policy_cfg = policy_config.llm_policy_cfg - - # [PRIORZERO-NEW] History buffer for each environment - # Format: {env_id: deque([(obs_text, action_text, reward), ...])} - self.history_buffers = defaultdict( - lambda: deque(maxlen=self.llm_policy_cfg.history_length) - ) - - # [PRIORZERO-NEW] Statistics for logging - self.llm_stats = { - 'total_calls': 0, - 'successful_calls': 0, - 'failed_calls': 0, - 'retry_count': 0, - 'total_latency': 0.0, - 'llm_prior_top1_match_count': 0, # How often LLM top-1 matches MCTS choice - } - - self._logger.info("✓ PriorZeroCollector initialized with vLLM engine") - self._logger.info(f" - History length: {self.llm_policy_cfg.history_length}") - self._logger.info(f" - Generate max length: {self.llm_policy_cfg.generate_max_len}") - - # [PRIORZERO-NEW] Use custom GameSegment - self.GameSegment = GameSegment - - async def _async_get_llm_prior( - self, - states: List[str], - request_ids: List[str], - histories: Optional[List[List[Tuple[str, str, float]]]] = None, - max_retries: int = 3, - timeout: float = 30.0 - ) -> List[Any]: - """ - [PRIORZERO-NEW] - Async call to LLM to get action ranking priors. - - Args: - states: List of current observation texts - request_ids: List of unique request IDs for tracking - histories: Optional list of history tuples for each state - max_retries: Maximum number of retries on failure - timeout: Timeout in seconds for each request - - Returns: - llm_outputs: List of vLLM output objects - """ - # [FIX] Check if vLLM engine is available - if self.vllm_engine is None: - self._logger.info("INFO: vLLM engine not available, skipping LLM prior") - return [None] * len(states) - - from priorzero_policy import build_llm_prompt - - # Build prompts - prompts = [] - for i, state in enumerate(states): - history = histories[i] if histories is not None else None - - # Build instruction using the helper function from policy - instruction = build_llm_prompt( - current_obs=state, - history=history, - use_cot=self.llm_policy_cfg.use_cot - ) - - # Apply chat template if policy has tokenizer - if hasattr(self._policy, 'llm_tokenizer'): - prompt = self._policy.llm_tokenizer.apply_chat_template( - [{"role": "user", "content": instruction}], - tokenize=False, - add_generation_prompt=True - ) - else: - prompt = instruction - - # [FIX] Ensure prompt is a string - if prompt is None: - self._logger.error(f"[ERROR] Prompt {i} is None! Instruction was: {instruction[:100] if instruction else 'None'}") - prompt = "" # Fallback to empty string - elif not isinstance(prompt, str): - self._logger.error(f"[ERROR] Prompt {i} is not a string! Type: {type(prompt)}, Value: {prompt}") - prompt = str(prompt) # Force conversion to string - - prompts.append(prompt) - - # Configure sampling parameters - sampling_params = SamplingParams( - temperature=1.0, - top_p=1.0, - max_tokens=self.llm_policy_cfg.generate_max_len, - skip_special_tokens=False, - ) - - # Retry logic - for attempt in range(max_retries): - try: - start_time = time.time() - - # [DEBUG] Log prompts and parameters before generation - if self.debug_mode and attempt == 0: - self._logger.info(f"[DEBUG] Sending {len(prompts)} prompts to vLLM engine") - for i, prompt in enumerate(prompts[:2]): # Show first 2 prompts - self._logger.info(f"[DEBUG] Prompt {i} (len={len(prompt)}): {prompt[:200]}...") - self._logger.info(f"[DEBUG] Sampling params: temp={sampling_params.temperature}, max_tokens={sampling_params.max_tokens}, top_p={sampling_params.top_p}") - self._logger.info(f"[DEBUG] Request IDs: {request_ids[:2]}...") - - # [FIX] vLLM V1 generate() takes single prompt, not list - # Create generators for each prompt individually - generators = [] - for i, (prompt, req_id) in enumerate(zip(prompts, request_ids)): - gen = self.vllm_engine.generate( - prompt, # Single prompt string - sampling_params, - req_id # Single request_id string - ) - generators.append((i, gen)) - - # Collect results - llm_outputs = [None] * len(prompts) - - try: - # Collect all results concurrently - async def collect_from_generator(idx, gen): - """Collect final result from a generator""" - final_result = None - async for result in gen: - final_result = result - # Check timeout - if time.time() - start_time > timeout: - raise asyncio.TimeoutError(f"LLM generation timeout after {timeout}s") - return idx, final_result - - # Gather all results concurrently - tasks = [collect_from_generator(idx, gen) for idx, gen in generators] - results = await asyncio.gather(*tasks, return_exceptions=True) - - # Process results - for result in results: - if isinstance(result, Exception): - raise result - idx, output = result - llm_outputs[idx] = output - - except asyncio.TimeoutError: - self._logger.warning(f"⚠ LLM generation timeout after {timeout}s (attempt {attempt+1}/{max_retries})") - if attempt < max_retries - 1: - self.llm_stats['retry_count'] += 1 - continue - else: - # On final timeout, return None for all - self.llm_stats['failed_calls'] += len(prompts) - return [None] * len(prompts) - - # Check if all outputs were received - if None in llm_outputs: - missing_count = llm_outputs.count(None) - self._logger.warning(f"⚠ {missing_count}/{len(prompts)} LLM outputs missing (attempt {attempt+1}/{max_retries})") - if attempt < max_retries - 1: - self.llm_stats['retry_count'] += 1 - continue - - # Success - elapsed = time.time() - start_time - self.llm_stats['total_calls'] += len(prompts) - self.llm_stats['successful_calls'] += len([o for o in llm_outputs if o is not None]) - self.llm_stats['failed_calls'] += len([o for o in llm_outputs if o is None]) - self.llm_stats['total_latency'] += elapsed - - self._logger.debug(f"✓ LLM generation completed in {elapsed:.2f}s ({len(prompts)} prompts)") - - # [DEBUG] Log detailed LLM outputs if debug mode is enabled - if self.debug_mode: - for i, (prompt, output) in enumerate(zip(prompts, llm_outputs)): - if output is not None: - output_text = output.outputs[0].text if output.outputs else "[No output]" - self._logger.info(f"[DEBUG] Env {i} - Prompt: {prompt[:100]}... -> LLM Output: {output_text[:100]}...") - else: - self._logger.warning(f"[DEBUG] Env {i} - LLM output is None") - - return llm_outputs - - except Exception as e: - import traceback - error_msg = f"{type(e).__name__}: {str(e)}" if str(e) else type(e).__name__ - error_trace = traceback.format_exc() - - # [FIX] Always log the full traceback on first attempt or in debug mode - if attempt == 0 or self.debug_mode: - self._logger.error(f"✗ LLM generation error (attempt {attempt+1}/{max_retries}): {error_msg}") - self._logger.error(f"Full traceback:\n{error_trace}") - else: - self._logger.error(f"✗ LLM generation error (attempt {attempt+1}/{max_retries}): {error_msg}") - - if attempt < max_retries - 1: - self.llm_stats['retry_count'] += 1 - await asyncio.sleep(0.5) # Brief pause before retry - else: - # Final failure - self._logger.error(f"✗ LLM generation failed after {max_retries} attempts. Last error: {error_msg}") - self._logger.error(f"Final traceback:\n{error_trace}") - self.llm_stats['failed_calls'] += len(prompts) - return [None] * len(prompts) - - return [None] * len(prompts) - - async def collect( - self, - num_segments: Optional[int] = None, - train_iter: int = 0, - policy_kwargs: Optional[dict] = None, - collect_with_pure_policy: bool = False - ) -> List[Any]: - """ - [PRIORZERO-MODIFIED] - Collect game segments with LLM-guided MCTS. - - Main changes from parent: - 1. Extract text observations from environment - 2. Async call to LLM to get action priors - 3. Pass LLM priors to policy forward pass - 4. Update history buffers after each step - - Args: - num_segments: Number of segments to collect - train_iter: Current training iteration - policy_kwargs: Additional kwargs for policy - collect_with_pure_policy: Whether to use pure policy without MCTS - - Returns: - return_data: List containing [game_segments, metadata] - """ - if num_segments is None: - if self._default_num_segments is None: - raise RuntimeError("Please specify num_segments for collection.") - else: - num_segments = self._default_num_segments - - assert num_segments == self._env_num, \ - f"num_segments({num_segments}) must equal env_num({self._env_num})" - - if policy_kwargs is None: - policy_kwargs = {} - - temperature = policy_kwargs.get('temperature', 1.0) - epsilon = policy_kwargs.get('epsilon', 0.0) - - # ================================================================== - # Initialization - # ================================================================== - collected_episode = 0 - collected_step = 0 - env_nums = self._env_num - init_obs = self._env.ready_obs - - # Wait for all environments to be ready - retry_waiting_time = 0.05 - while len(init_obs.keys()) != env_nums: - self._logger.info(f'Waiting for all environments to reset. Ready: {list(init_obs.keys())}') - time.sleep(retry_waiting_time) - init_obs = self._env.ready_obs - - # Initialize state tracking - for env_id in range(env_nums): - if env_id in init_obs: - self.action_mask_dict[env_id] = to_ndarray(init_obs[env_id]['action_mask']) - self.to_play_dict[env_id] = to_ndarray(init_obs[env_id]['to_play']) - self.timestep_dict[env_id] = to_ndarray(init_obs[env_id].get('timestep', -1)) - - # Initialize game segments - game_segments = [ - GameSegment( - self._env.action_space, - game_segment_length=self.policy_config.game_segment_length, - config=self.policy_config, - task_id=self.task_id - ) for _ in range(env_nums) - ] - - # Initialize observation stacks - observation_window_stack = [ - deque(maxlen=self.policy_config.model.frame_stack_num) - for _ in range(env_nums) - ] - for env_id in range(env_nums): - initial_frames = [ - to_ndarray(init_obs[env_id]['observation']) - for _ in range(self.policy_config.model.frame_stack_num) - ] - observation_window_stack[env_id].extend(initial_frames) - game_segments[env_id].reset(observation_window_stack[env_id]) - - # Priority calculation lists - search_values_lst = [[] for _ in range(env_nums)] - pred_values_lst = [[] for _ in range(env_nums)] - - # Logging variables - eps_steps_lst = np.zeros(env_nums) - visit_entropies_lst = np.zeros(env_nums) - - if collect_with_pure_policy: - temp_visit_list = [0.0 for _ in range(self._env.action_space.n)] - - # ================================================================== - # Main Collection Loop - # ================================================================== - while True: - with self._timer: - # Get ready environments - obs = self._env.ready_obs - ready_env_id = set(obs.keys()) - - if len(ready_env_id) < self._env_num: - self._logger.debug(f'Only {len(ready_env_id)}/{self._env_num} envs ready') - - # Prepare stacked observations for world model - stack_obs_dict = { - env_id: game_segments[env_id].get_obs() - for env_id in ready_env_id - } - stack_obs_list = [stack_obs_dict[env_id] for env_id in sorted(list(ready_env_id))] - - # Prepare action masks and other info - action_mask = [self.action_mask_dict[env_id] for env_id in sorted(list(ready_env_id))] - to_play = [self.to_play_dict[env_id] for env_id in sorted(list(ready_env_id))] - timestep = [self.timestep_dict[env_id] for env_id in sorted(list(ready_env_id))] - - # Convert to tensors - stack_obs_array = to_ndarray(stack_obs_list) - stack_obs_tensor = prepare_observation( - stack_obs_array, - self.policy_config.model.model_type - ) - stack_obs_tensor = torch.from_numpy(stack_obs_tensor).to(self.policy_config.device) - - # ============================================================== - # [PRIORZERO-NEW] Get LLM Priors - # ============================================================== - if not collect_with_pure_policy: - # Extract text observations and valid actions - raw_obs_list = [] - histories_list = [] - valid_actions_list = [] # [PRIORZERO] Store valid actions for each env - for env_id in sorted(list(ready_env_id)): - # Extract raw text - raw_obs_text = extract_raw_obs_text(obs[env_id]) - raw_obs_list.append(raw_obs_text) - - # Get history for this environment - history = list(self.history_buffers[env_id]) - histories_list.append(history) - - # [PRIORZERO] Extract valid actions from observation - valid_actions = obs[env_id].get('valid_actions', []) - valid_actions_list.append(valid_actions) - - # Generate request IDs - request_ids = [ - f"collect_{train_iter}_{i}" - for i in range(len(raw_obs_list)) - ] - - # Async call to LLM - llm_outputs = await self._async_get_llm_prior( - raw_obs_list, - request_ids, - histories_list - ) - - # Add to policy kwargs - policy_kwargs['llm_prior_outputs'] = llm_outputs - policy_kwargs['valid_actions_list'] = valid_actions_list # [PRIORZERO] Pass valid actions - else: - policy_kwargs['llm_prior_outputs'] = None - policy_kwargs['valid_actions_list'] = None - - # ============================================================== - # Policy Forward Pass - # ============================================================== - policy_args = (stack_obs_tensor, action_mask, temperature, to_play, epsilon) - policy_kwargs_forward = { - 'ready_env_id': sorted(list(ready_env_id)), - 'timestep': timestep, - 'llm_prior_outputs': policy_kwargs.get('llm_prior_outputs'), - 'valid_actions_list': policy_kwargs.get('valid_actions_list') # [PRIORZERO] Pass valid actions - } - - if self.task_id is not None: - policy_kwargs_forward['task_id'] = self.task_id - - policy_output = self._policy.forward(*policy_args, **policy_kwargs_forward) - - # Extract outputs - actions_with_env_id = {k: v['action'] for k, v in policy_output.items()} - value_dict_with_env_id = {k: v['searched_value'] for k, v in policy_output.items()} - pred_value_dict_with_env_id = {k: v['predicted_value'] for k, v in policy_output.items()} - - if not collect_with_pure_policy: - distributions_dict_with_env_id = { - k: v['visit_count_distributions'] for k, v in policy_output.items() - } - visit_entropy_dict_with_env_id = { - k: v['visit_count_distribution_entropy'] for k, v in policy_output.items() - } - - actions: Dict[int, Any] = { - env_id: actions_with_env_id.pop(env_id) - for env_id in ready_env_id - } - - # ============================================================== - # Step Environments - # ============================================================== - timesteps = self._env.step(actions) - - # [DEBUG] Log actions taken if debug mode is enabled - if self.debug_mode: - for env_id, action in actions.items(): - self._logger.info(f"[DEBUG] Env {env_id} - Action taken: {action}") - - interaction_duration = self._timer.value / len(timesteps) - - # ================================================================== - # Process Environment Responses - # ================================================================== - for env_id, episode_timestep in timesteps.items(): - with self._timer: - # Handle abnormal timesteps - if episode_timestep.info.get('abnormal', False): - self._env.reset({env_id: None}) - self._policy.reset([env_id]) - self._reset_stat(env_id) - self._logger.info(f'⚠ Env {env_id} had abnormal step: {episode_timestep.info}') - continue - - obs_new, reward, done, info = ( - episode_timestep.obs, - episode_timestep.reward, - episode_timestep.done, - episode_timestep.info - ) - - # [DEBUG] Log observation and reward if debug mode is enabled - if self.debug_mode: - raw_obs_preview = extract_raw_obs_text(obs_new)[:150] - self._logger.info(f"[DEBUG] Env {env_id} - Obs: {raw_obs_preview}... | Reward: {reward} | Done: {done}") - - # Store search statistics - if collect_with_pure_policy: - game_segments[env_id].store_search_stats(temp_visit_list, 0) - else: - game_segments[env_id].store_search_stats( - distributions_dict_with_env_id[env_id], - value_dict_with_env_id[env_id] - ) - - # Append transition to game segment - # [PRIORZERO-FIX] Extract and pass raw_obs_text to GameSegment - raw_obs_text_for_segment = extract_raw_obs_text(obs_new) - - game_segments[env_id].append( - actions[env_id], - to_ndarray(obs_new['observation']), - reward, - self.action_mask_dict[env_id], - self.to_play_dict[env_id], - timestep=to_ndarray(obs_new.get('timestep', -1)), - raw_obs_text=raw_obs_text_for_segment - ) - - # =========================================================== - # [PRIORZERO-NEW] Update History Buffer - # =========================================================== - raw_obs_text = extract_raw_obs_text(obs[env_id]) - # [PRIORZERO] Use dynamic action mapping if available - dynamic_action_inv_map = policy_output.get(env_id, {}).get('dynamic_action_inv_map', None) - if dynamic_action_inv_map is not None: - action_text = dynamic_action_inv_map.get(actions[env_id], f"action_{actions[env_id]}") - else: - # Fallback to static mapping - action_text = getattr(self._policy, 'action_inv_map', {}).get( - actions[env_id], - f"action_{actions[env_id]}" - ) - self.history_buffers[env_id].append((raw_obs_text, action_text, float(reward))) - - # Update state - self.action_mask_dict[env_id] = to_ndarray(obs_new['action_mask']) - self.to_play_dict[env_id] = to_ndarray(obs_new['to_play']) - self.timestep_dict[env_id] = to_ndarray(obs_new.get('timestep', -1)) - self.dones[env_id] = False if self.policy_config.ignore_done else done - - if not collect_with_pure_policy: - visit_entropies_lst[env_id] += visit_entropy_dict_with_env_id[env_id] - - eps_steps_lst[env_id] += 1 - - # Reset policy if needed (for UniZero) - if self._policy.get_attribute('cfg').type in ['unizero', 'sampled_unizero', 'priorzero']: - self._policy.reset( - env_id=env_id, - current_steps=eps_steps_lst[env_id], - reset_init_data=False - ) - - # Store values for priority calculation - if self.policy_config.use_priority: - pred_values_lst[env_id].append(pred_value_dict_with_env_id[env_id]) - search_values_lst[env_id].append(value_dict_with_env_id[env_id]) - - # Update observation window - observation_window_stack[env_id].append(to_ndarray(obs_new['observation'])) - - # =========================================================== - # Save Full Game Segment - # =========================================================== - if game_segments[env_id].is_full(): - if self.last_game_segments[env_id] is not None: - self.pad_and_save_last_trajectory( - env_id, - self.last_game_segments, - self.last_game_priorities, - game_segments, - self.dones - ) - - # Calculate priorities - priorities = self._compute_priorities( - env_id, - pred_values_lst, - search_values_lst - ) - pred_values_lst[env_id], search_values_lst[env_id] = [], [] - - # Save segment - self.last_game_segments[env_id] = game_segments[env_id] - self.last_game_priorities[env_id] = priorities - - # Create new segment - game_segments[env_id] = GameSegment( - self._env.action_space, - game_segment_length=self.policy_config.game_segment_length, - config=self.policy_config, - task_id=self.task_id - ) - game_segments[env_id].reset(observation_window_stack[env_id]) - - self._env_info[env_id]['step'] += 1 - collected_step += 1 - - self._env_info[env_id]['time'] += self._timer.value + interaction_duration - - # ============================================================== - # Episode Done - # ============================================================== - if episode_timestep.done: - self._logger.info(f'======== Env {env_id} episode finished! ========') - self._total_episode_count += 1 - - # Logging - info_log = { - 'reward': episode_timestep.info['eval_episode_return'], - 'time': self._env_info[env_id]['time'], - 'step': self._env_info[env_id]['step'], - } - if not collect_with_pure_policy: - info_log['visit_entropy'] = ( - visit_entropies_lst[env_id] / eps_steps_lst[env_id] - if eps_steps_lst[env_id] > 0 else 0 - ) - - collected_episode += 1 - self._episode_info.append(info_log) - - # Save remaining segments - if self.last_game_segments[env_id] is not None: - self.pad_and_save_last_trajectory( - env_id, - self.last_game_segments, - self.last_game_priorities, - game_segments, - self.dones - ) - - priorities = self._compute_priorities( - env_id, - pred_values_lst, - search_values_lst - ) - - game_segments[env_id].game_segment_to_array() - if len(game_segments[env_id].reward_segment) > 0: - self.game_segment_pool.append(( - game_segments[env_id], - priorities, - self.dones[env_id] - )) - - # Reset - pred_values_lst[env_id], search_values_lst[env_id] = [], [] - eps_steps_lst[env_id], visit_entropies_lst[env_id] = 0, 0 - - self._policy.reset([env_id], task_id=self.task_id) - self._reset_stat(env_id) - - # Clear history buffer for this environment - self.history_buffers[env_id].clear() - - # Re-initialize game segment - game_segments[env_id] = GameSegment( - self._env.action_space, - game_segment_length=self.policy_config.game_segment_length, - config=self.policy_config, - task_id=self.task_id - ) - game_segments[env_id].reset(observation_window_stack[env_id]) - - # ================================================================== - # Check if Enough Segments Collected - # ================================================================== - if len(self.game_segment_pool) >= self._default_num_segments: - self._logger.info( - f'✓ Collected {len(self.game_segment_pool)} segments ' - f'(target: {self._default_num_segments})' - ) - - # Format return data - return_data = [ - [self.game_segment_pool[i][0] for i in range(len(self.game_segment_pool))], - [ - { - 'priorities': self.game_segment_pool[i][1], - 'done': self.game_segment_pool[i][2], - 'unroll_plus_td_steps': self.unroll_plus_td_steps - } - for i in range(len(self.game_segment_pool)) - ] - ] - self.game_segment_pool.clear() - break - - # ================================================================== - # Final Logging - # ================================================================== - collected_duration = sum([d['time'] for d in self._episode_info]) - - self._total_envstep_count += collected_step - self._total_episode_count += collected_episode - self._total_duration += collected_duration - - self._output_log(train_iter) - - # [PRIORZERO-NEW] Log LLM statistics - if self.llm_stats['total_calls'] > 0: - avg_latency = self.llm_stats['total_latency'] / self.llm_stats['total_calls'] - success_rate = self.llm_stats['successful_calls'] / self.llm_stats['total_calls'] - - self._logger.info( - f"📊 LLM Prior Statistics:\n" - f" - Total calls: {self.llm_stats['total_calls']}\n" - f" - Success rate: {success_rate*100:.1f}%\n" - f" - Avg latency: {avg_latency:.3f}s\n" - f" - Retry count: {self.llm_stats['retry_count']}" - ) - - return return_data - - def _output_log(self, train_iter: int) -> None: - """ - [INHERITED] - Log collection statistics (inherited from parent). - """ - super()._output_log(train_iter) diff --git a/zoo/jericho/priorzero/priorzero_collector_unified.py b/zoo/jericho/priorzero/priorzero_collector_unified.py new file mode 100644 index 000000000..1c3d7cb8b --- /dev/null +++ b/zoo/jericho/priorzero/priorzero_collector_unified.py @@ -0,0 +1,740 @@ +""" +Unified PriorZero Collector supporting both LLM and VL priors + +This collector uses a unified prior_generator interface to support: +- Text input with LLM prior (Jericho games) +- Image input with VL prior (Atari games) +""" +import asyncio +import logging +import sys +import time + +from collections import deque, defaultdict +from pathlib import Path +from typing import Optional, Any, List, Dict, Tuple + +import numpy as np +import torch +from ding.envs import BaseEnvManager +from ding.torch_utils import to_ndarray +from ding.utils import build_logger, EasyTimer, SERIAL_COLLECTOR_REGISTRY, allreduce_data +from vllm import SamplingParams +import os + +# Import from local LightZero +from lzero.worker.muzero_segment_collector import MuZeroSegmentCollector as OriginalCollector +from lzero.mcts.utils import prepare_observation +from game_segment_priorzero import GameSegment + + +# ============================================================================== +# Helper Functions +# ============================================================================== + +def extract_raw_obs_text(obs_dict: Dict[str, Any]) -> str: + """Extract text observation from environment observation dictionary.""" + if 'raw_obs_text' in obs_dict: + return str(obs_dict['raw_obs_text']) + if 'raw_obs' in obs_dict: + return str(obs_dict['raw_obs']) + if 'text' in obs_dict: + return str(obs_dict['text']) + if 'observation_str' in obs_dict: + return str(obs_dict['observation_str']) + if 'observation' in obs_dict: + obs = obs_dict['observation'] + if isinstance(obs, str): + return obs + elif isinstance(obs, (list, np.ndarray)): + return f"[Observation vector of shape {np.array(obs).shape}]" + return str(obs_dict) + + +def extract_raw_obs_image(obs_dict: Dict[str, Any]) -> np.ndarray: + """Extract image observation from environment observation dictionary.""" + if 'observation' in obs_dict: + obs = obs_dict['observation'] + if isinstance(obs, np.ndarray): + # Assume image format (H, W, C) or (C, H, W) + return obs + raise ValueError(f"Cannot extract image from observation: {obs_dict.keys()}") + + +# ============================================================================== +# Unified PriorZero Collector Class +# ============================================================================== + +@SERIAL_COLLECTOR_REGISTRY.register('priorzero_segment', force_overwrite=True) +class PriorZeroCollector(OriginalCollector): + """ + Unified PriorZero Collector supporting both LLM and VL priors. + + Features: + - Unified prior_generator interface (supports LLM and VL) + - History buffer for each environment + - Automatic detection of observation type (text vs image) + - Backward compatible with existing LLM-based implementation + """ + + def __init__( + self, + policy_config: Dict, + llm_config: Dict, # Can be LLM or VL config + data_processor=None, # Backward compatibility + prior_generator=None, # NEW: Unified prior generator + prof=None, + obs_type: str = 'text', # NEW: 'text' or 'image' + env_id: str = None, # NEW: Environment ID for action mapping + **kwargs + ): + """ + Initialize Unified PriorZeroCollector. + + Args: + policy_config: Policy configuration + llm_config: LLM/VL configuration + data_processor: DataProcessor (for backward compatibility) + prior_generator: Unified PriorGenerator instance (NEW) + prof: Profiler + obs_type: Observation type ('text' or 'image') + env_id: Environment ID (e.g., 'PongNoFrameskip-v4') + **kwargs: Additional arguments for parent class + """ + kwargs['policy_config'] = policy_config + + super().__init__(**kwargs) + + self.data_processor = data_processor + self.prior_generator = prior_generator # NEW: Unified interface + self.prof = prof + self.llm_cfg = llm_config + self.obs_type = obs_type # NEW: Track observation type + self.env_id = env_id or 'PongNoFrameskip-v4' # NEW: Store env_id + + # History buffers + history_length = getattr(llm_config, 'history_length', 5) + self.history_buffers = defaultdict(lambda: deque(maxlen=history_length)) + self.llm_prior_temperature = getattr(llm_config, 'llm_prior_temperature', 1.0) + + # Logging + prior_type = "VL" if obs_type == 'image' else "LLM" + self._logger.info(f"✓ PriorZeroCollector initialized with {prior_type} prior") + self._logger.info(f" - Observation type: {obs_type}") + if obs_type == 'image': + self._logger.info(f" - Environment: {self.env_id}") + self._logger.info(f" - History length: {history_length}") + self._logger.info(f" - Prior generator: {type(prior_generator).__name__ if prior_generator else 'None'}") + + # First-call validation flag + self._first_collect_logged = False + + def _get_prior_from_generator( + self, + observations: List[Any], + valid_actions_list: List[List[str]], + histories_list: List[List], + ) -> Tuple[List[np.ndarray], List[np.ndarray], List[Any]]: + """ + Get action priors using the unified prior_generator interface. + + Args: + observations: List of observations (text strings or image arrays) + valid_actions_list: List of valid action lists + histories_list: List of history buffers + + Returns: + Tuple of (prior_per_seq, prior_per_tok, cot_prefixes) + """ + if self.prior_generator is None: + # Fallback: uniform prior + num_envs = len(observations) + prior_per_seq = [] + for actions in valid_actions_list: + uniform_prior = np.ones(len(actions)) / len(actions) + prior_per_seq.append(uniform_prior) + prior_per_tok = [None] * num_envs + cot_prefixes = [None] * num_envs + return prior_per_seq, prior_per_tok, cot_prefixes + + # Use unified prior generator + prior_results = self.prior_generator.batch_generate_prior( + observations=observations, + action_candidates_list=valid_actions_list, + histories=histories_list, + temperature=self.llm_prior_temperature, + ) + + # Extract results + prior_per_seq = [result['action_probs'] for result in prior_results] + # VL path: action_logits is np.ndarray of shape (num_actions,) with per-action log-probs. + # The VL datafactory (_make_vl_train_samples) expects this numpy format and spreads + # the chosen action's logprob uniformly across target tokens for PPO. + # NOTE: Do NOT convert to dict here — the game_buffer's _is_llm_text_mode guard + # uses isinstance(..., dict) to distinguish LLM text mode from VL image mode. + prior_per_tok = [result.get('action_logits', None) for result in prior_results] + cot_prefixes = [result.get('raw_output', None) for result in prior_results] + + return prior_per_seq, prior_per_tok, cot_prefixes + + def _get_prior_legacy( + self, + raw_obs_list: List[str], + valid_actions_list: List[List[str]], + histories_list: List[List], + ) -> Tuple[List[np.ndarray], List[np.ndarray], List[Any]]: + """ + Legacy method using data_processor (for backward compatibility). + + This is the original implementation for LLM-based priors. + """ + if self.data_processor is None: + raise ValueError("data_processor is None. Cannot use legacy prior generation.") + + llm_prior_per_seq, llm_prior_per_tok, cot_prefixes = self.data_processor.get_llm_prior( + states=raw_obs_list, + valid_actions_list=valid_actions_list, + histories=histories_list, + return_cot=True + ) + + return llm_prior_per_seq, llm_prior_per_tok, cot_prefixes + + def collect( + self, + num_segments: Optional[int] = None, + train_iter: int = 0, + policy_kwargs: Optional[dict] = None, + collect_with_pure_policy: bool = False + ) -> List[Any]: + """ + Collect game segments with prior-guided MCTS. + + Supports both LLM (text) and VL (image) priors through unified interface. + + Args: + num_segments: Number of segments to collect + train_iter: Current training iteration + policy_kwargs: Additional kwargs for policy + collect_with_pure_policy: Whether to use pure policy without MCTS + + Returns: + return_data: List containing [game_segments, metadata] + """ + if num_segments is None: + if self._default_num_segments is None: + raise RuntimeError("Please specify num_segments for collection.") + else: + num_segments = self._default_num_segments + + assert num_segments == self._env_num, \ + f"num_segments({num_segments}) must equal env_num({self._env_num})" + + if policy_kwargs is None: + policy_kwargs = {} + + temperature = policy_kwargs.get('temperature', 1.0) + epsilon = policy_kwargs.get('epsilon', 0.0) + + collected_episode = 0 + collected_step = 0 + llm_prior_entropy = [[] for _ in range(self._env_num)] + env_nums = self._env_num + init_obs = self._env.ready_obs + + retry_waiting_time = 0.05 + while len(init_obs.keys()) != env_nums: + self._logger.info(f'Waiting for all environments to reset. Ready: {list(init_obs.keys())}') + time.sleep(retry_waiting_time) + init_obs = self._env.ready_obs + + for env_id in range(env_nums): + if env_id in init_obs: + self.action_mask_dict[env_id] = to_ndarray(init_obs[env_id]['action_mask']) + self.to_play_dict[env_id] = to_ndarray(init_obs[env_id]['to_play']) + self.timestep_dict[env_id] = to_ndarray(init_obs[env_id].get('timestep', -1)) + + last_game_segments = [None for _ in range(env_nums)] + last_game_priorities = [None for _ in range(env_nums)] + game_segments = [ + GameSegment( + self._env.action_space, + game_segment_length=self.policy_config.game_segment_length, + config=self.policy_config, + task_id=self.task_id + ) for _ in range(env_nums) + ] + + observation_window_stack = [ + deque(maxlen=self.policy_config.model.frame_stack_num) + for _ in range(env_nums) + ] + for env_id in range(env_nums): + initial_frames = [ + to_ndarray(init_obs[env_id]['observation']) + for _ in range(self.policy_config.model.frame_stack_num) + ] + observation_window_stack[env_id].extend(initial_frames) + + # Extract initial raw observation (text or image) + if self.obs_type == 'text': + init_raw_obs = extract_raw_obs_text(init_obs[env_id]) + else: + init_raw_obs = extract_raw_obs_image(init_obs[env_id]) + + game_segments[env_id].reset( + observation_window_stack[env_id], + init_raw_obs=init_raw_obs, + init_history_obs=list(self.history_buffers[env_id]) + ) + + search_values_lst = [[] for _ in range(env_nums)] + pred_values_lst = [[] for _ in range(env_nums)] + + eps_steps_lst = np.zeros(env_nums) + visit_entropies_lst = np.zeros(env_nums) + + if collect_with_pure_policy: + temp_visit_list = [0.0 for _ in range(self._env.action_space.n)] + + while True: + with self._timer: + obs = self._env.ready_obs + ready_env_id = set(obs.keys()) + + if len(ready_env_id) < self._env_num: + self._logger.debug(f'Only {len(ready_env_id)}/{self._env_num} envs ready') + + stack_obs_dict = { + env_id: game_segments[env_id].get_obs() + for env_id in ready_env_id + } + stack_obs_list = [stack_obs_dict[env_id] for env_id in sorted(list(ready_env_id))] + + action_mask = [self.action_mask_dict[env_id] for env_id in sorted(list(ready_env_id))] + to_play = [self.to_play_dict[env_id] for env_id in sorted(list(ready_env_id))] + timestep = [self.timestep_dict[env_id] for env_id in sorted(list(ready_env_id))] + + # Convert to tensors + stack_obs_array = to_ndarray(stack_obs_list) + stack_obs_tensor = prepare_observation( + stack_obs_array, + self.policy_config.model.model_type + ) + stack_obs_tensor = torch.from_numpy(stack_obs_tensor).to(self.policy_config.device).float() + + if collect_with_pure_policy: + continue + else: + # =========================================================== + # [UNIFIED] Extract observations and get priors + # =========================================================== + observations_list = [] + histories_list = [] + valid_actions_list = [] + + for env_id in sorted(list(ready_env_id)): + # Extract observation based on type + if self.obs_type == 'text': + raw_obs = extract_raw_obs_text(obs[env_id]) + else: # image + raw_obs = extract_raw_obs_image(obs[env_id]) + + observations_list.append(raw_obs) + histories_list.append(list(self.history_buffers[env_id])) + + # Get valid actions + # For text games: use valid_actions from obs + # For Atari: convert integer indices to semantic action names + valid_actions = obs[env_id].get('valid_actions', []) + if len(valid_actions) == 0 and self.obs_type == 'image': + # Atari: convert integer action indices to semantic names + from zoo.jericho.priorzero.atari_action_meanings import get_action_meanings + action_space_size = self.policy_config.model.action_space_size + action_meanings = get_action_meanings(self.env_id, action_space_size) + # Use semantic names instead of integers + valid_actions = [action_meanings[i] for i in range(action_space_size)] + valid_actions_list.append(valid_actions) + + # First-call validation logging for image data flow + if not self._first_collect_logged and self.obs_type == 'image' and len(observations_list) > 0: + self._first_collect_logged = True + obs_sample = observations_list[0] + if isinstance(obs_sample, np.ndarray): + self._logger.info( + f"[Collector Validation] === FIRST COLLECT IMAGE CHECK ===\n" + f" Image shape: {obs_sample.shape}, dtype: {obs_sample.dtype}, " + f"min: {obs_sample.min()}, max: {obs_sample.max()}\n" + f" Num envs: {len(observations_list)}\n" + f" Actions: {valid_actions_list[0]}\n" + f"[Collector Validation] === END CHECK ===" + ) + + # Get priors using unified interface + with self.prof.block("collect_step_get_prior", rank=self._rank): + if self.prior_generator is not None: + # NEW: Use unified prior generator + llm_prior_per_seq, llm_prior_per_tok, cot_prefixes = self._get_prior_from_generator( + observations=observations_list, + valid_actions_list=valid_actions_list, + histories_list=histories_list, + ) + elif self.data_processor is not None: + # LEGACY: Use data_processor (backward compatibility) + llm_prior_per_seq, llm_prior_per_tok, cot_prefixes = self._get_prior_legacy( + raw_obs_list=observations_list, + valid_actions_list=valid_actions_list, + histories_list=histories_list, + ) + else: + # Fallback: uniform prior + llm_prior_per_seq = [ + np.ones(len(actions)) / len(actions) + for actions in valid_actions_list + ] + llm_prior_per_tok = [None] * len(observations_list) + cot_prefixes = [None] * len(observations_list) + + # Apply temperature scaling + for env_id, llm_prior in enumerate(llm_prior_per_seq): + scaled_llm_prior = self.apply_temperature_scaling(llm_prior, return_logprobs=True) + llm_prior_per_seq[env_id] = scaled_llm_prior + + policy_kwargs_forward = { + 'llm_prior_logprob': llm_prior_per_seq, + 'valid_actions_list': valid_actions_list, + } + + if self.task_id is not None: + policy_kwargs_forward['task_id'] = self.task_id + + with self.prof.block("collect_step_forward", rank=self._rank): + policy_output = self._policy.forward( + data=stack_obs_tensor, + action_mask=action_mask, + temperature=temperature, + to_play=to_play, + epsilon=epsilon, + ready_env_id=sorted(list(ready_env_id)), + timestep=timestep, + **policy_kwargs_forward + ) + + # Extract outputs + actions_with_env_id = {k: v['action'] for k, v in policy_output.items()} + value_dict_with_env_id = {k: v['searched_value'] for k, v in policy_output.items()} + pred_value_dict_with_env_id = {k: v['predicted_value'] for k, v in policy_output.items()} + + if not collect_with_pure_policy: + distributions_dict_with_env_id = { + k: v['visit_count_distributions'] for k, v in policy_output.items() + } + visit_entropy_dict_with_env_id = { + k: v['visit_count_distribution_entropy'] for k, v in policy_output.items() + } + + actions: Dict[int, Any] = { + env_id: actions_with_env_id.pop(env_id) + for env_id in ready_env_id + } + + with self.prof.block("collect_step", rank=self._rank): + timesteps = self._env.step(actions) + + interaction_duration = self._timer.value / len(timesteps) + + for env_id, episode_timestep in timesteps.items(): + with self._timer: + # Handle abnormal timesteps + if episode_timestep.info.get('abnormal', False): + self._env.reset({env_id: None}) + self._policy.reset([env_id]) + self._reset_stat(env_id) + self._logger.info(f'⚠ Env {env_id} had abnormal step: {episode_timestep.info}') + continue + + obs_new, reward, done, info = ( + episode_timestep.obs, + episode_timestep.reward, + episode_timestep.done, + episode_timestep.info + ) + + game_segments[env_id].store_search_stats( + distributions_dict_with_env_id[env_id], + value_dict_with_env_id[env_id] + ) + + # =========================================================== + # [UNIFIED] Update History Buffer + # =========================================================== + if self.obs_type == 'text': + raw_obs = extract_raw_obs_text(obs[env_id]) + else: + raw_obs = extract_raw_obs_image(obs[env_id]) + + # Get action string + # For Atari: convert integer action index to semantic name + if self.obs_type == 'image': + from zoo.jericho.priorzero.atari_action_meanings import action_index_to_name + action_space_size = self.policy_config.model.action_space_size + action_str = action_index_to_name(self.env_id, actions[env_id], action_space_size) + elif env_id < len(valid_actions_list) and actions[env_id] < len(valid_actions_list[env_id]): + # Text games: use action name from valid_actions_list + action_str = valid_actions_list[env_id][actions[env_id]] + else: + # Fallback + action_str = info.get('action_str', str(actions[env_id])) + + # Use absolute timestep from environment, not relative episode step counter + abs_timestep = int(self.timestep_dict[env_id]) if int(self.timestep_dict[env_id]) >= 0 else int(eps_steps_lst[env_id]) + self.history_buffers[env_id].append((raw_obs, action_str, float(reward), abs_timestep)) + + # Append transition to game segment + game_segments[env_id].append( + actions[env_id], + to_ndarray(obs_new['observation']), + reward, + self.action_mask_dict[env_id], + self.to_play_dict[env_id], + raw_obs_text=raw_obs, + history_obs=list(self.history_buffers[env_id]), + llm_prior_per_tok=llm_prior_per_tok[env_id] if env_id < len(llm_prior_per_tok) else None, + cot_prefix=cot_prefixes[env_id] if env_id < len(cot_prefixes) else None, + llm_action=action_str + ) + + # Update statistics + self.action_mask_dict[env_id] = to_ndarray(obs_new['action_mask']) + self.to_play_dict[env_id] = to_ndarray(obs_new['to_play']) + self.timestep_dict[env_id] = to_ndarray(obs_new.get('timestep', -1)) + + observation_window_stack[env_id].append(to_ndarray(obs_new['observation'])) + + search_values_lst[env_id].append(value_dict_with_env_id[env_id]) + pred_values_lst[env_id].append(pred_value_dict_with_env_id[env_id]) + + if not collect_with_pure_policy: + visit_entropies_lst[env_id] += visit_entropy_dict_with_env_id[env_id] + + eps_steps_lst[env_id] += 1 + collected_step += 1 + + # Check if segment is complete + if game_segments[env_id].is_full(): + if last_game_segments[env_id] is not None: + self.pad_and_save_last_trajectory( + env_id, last_game_segments, last_game_priorities, + game_segments, done + ) + + last_game_segments[env_id] = game_segments[env_id] + last_game_priorities[env_id] = self._compute_priorities(game_segments[env_id]) + + # Create new segment + game_segments[env_id] = GameSegment( + self._env.action_space, + game_segment_length=self.policy_config.game_segment_length, + config=self.policy_config, + task_id=self.task_id + ) + + if self.obs_type == 'text': + current_raw_obs = extract_raw_obs_text(obs_new) + else: + current_raw_obs = extract_raw_obs_image(obs_new) + + game_segments[env_id].reset( + observation_window_stack[env_id], + init_raw_obs=current_raw_obs, + init_history_obs=list(self.history_buffers[env_id]) + ) + + # Handle episode end + if done: + self._env.reset({env_id: None}) + self._policy.reset([env_id]) + + # Save second-to-last segment (if exists) + if last_game_segments[env_id] is not None: + self.pad_and_save_last_trajectory( + env_id, last_game_segments, last_game_priorities, + game_segments, done + ) + + # Save the final segment of the episode + game_segments[env_id].game_segment_to_array() + if len(game_segments[env_id].reward_segment) > 0: + priorities = self._compute_priorities(game_segments[env_id]) + self.game_segment_pool.append((game_segments[env_id], priorities, done)) + + # Log episode statistics + collected_episode += 1 + episode_return = info.get('eval_episode_return', info.get('score', reward)) + self._logger.info( + f"Episode {collected_episode} | Env {env_id} | " + f"Steps: {eps_steps_lst[env_id]} | " + f"Reward: {episode_return:.2f}" + ) + + # Populate _episode_info for parent's _output_log() and TB logging + ep_info = { + 'reward': episode_return, + 'time': interaction_duration * eps_steps_lst[env_id], + 'step': int(eps_steps_lst[env_id]), + 'visit_entropy': visit_entropies_lst[env_id] / max(eps_steps_lst[env_id], 1), + } + self._episode_info.append(ep_info) + + # TB logging for episode metrics + if hasattr(self, '_tb_logger') and self._tb_logger is not None: + self._tb_logger.add_scalar('collect/episode_reward', episode_return, self._total_envstep_count + collected_step) + self._tb_logger.add_scalar('collect/episode_length', eps_steps_lst[env_id], self._total_envstep_count + collected_step) + + # Reset for next episode + eps_steps_lst[env_id] = 0 + visit_entropies_lst[env_id] = 0 + search_values_lst[env_id] = [] + pred_values_lst[env_id] = [] + self.history_buffers[env_id].clear() + + # Re-initialize game segment for next episode + init_obs = self._env.ready_obs + if env_id in init_obs: + game_segments[env_id] = GameSegment( + self._env.action_space, + game_segment_length=self.policy_config.game_segment_length, + config=self.policy_config, + task_id=self.task_id + ) + observation_window_stack[env_id] = deque(maxlen=self.policy_config.model.frame_stack_num) + initial_frames = [ + to_ndarray(init_obs[env_id]['observation']) + for _ in range(self.policy_config.model.frame_stack_num) + ] + observation_window_stack[env_id].extend(initial_frames) + + if self.obs_type == 'text': + init_raw_obs = extract_raw_obs_text(init_obs[env_id]) + else: + init_raw_obs = extract_raw_obs_image(init_obs[env_id]) + + game_segments[env_id].reset( + observation_window_stack[env_id], + init_raw_obs=init_raw_obs, + init_history_obs=list(self.history_buffers[env_id]) + ) + + self.action_mask_dict[env_id] = to_ndarray(init_obs[env_id]['action_mask']) + self.to_play_dict[env_id] = to_ndarray(init_obs[env_id]['to_play']) + + last_game_segments[env_id] = None + last_game_priorities[env_id] = None + + # Check if collection is complete + if collected_episode >= num_segments: + break + + # Update statistics that the parent class normally maintains + self._total_envstep_count += collected_step + self._total_episode_count += collected_episode + + # Call parent's _output_log to write standard TB metrics (collector_step/xxx) + # Only call when tb_logger is available (rank 0); parent _output_log has no None guard. + if self._tb_logger is not None: + self._output_log(train_iter) + + # Return collected data in the format expected by push_game_segments: + # [list_of_game_segments, list_of_meta_dicts] + return_data = [ + [seg for seg, _, _ in self.game_segment_pool], + [ + { + 'priorities': priorities, + 'done': done, + 'unroll_plus_td_steps': self.unroll_plus_td_steps, + } + for _, priorities, done in self.game_segment_pool + ] + ] + self.game_segment_pool = [] + + return return_data + + def pad_and_save_last_trajectory( + self, i: int, last_game_segments: List[GameSegment], last_game_priorities: List[np.ndarray], + game_segments: List[GameSegment], done: bool + ) -> None: + """Pad and save the last trajectory (same as original).""" + beg_index = self.policy_config.model.frame_stack_num + end_index = beg_index + self.policy_config.num_unroll_steps + self.policy_config.td_steps + + pad_obs_lst = game_segments[i].obs_segment[beg_index:end_index] + pad_raw_obs_lst = game_segments[i].raw_obs_segment[beg_index:end_index] + pad_history_obs_lst = game_segments[i].history_obs_segment[beg_index:end_index] + pad_llm_prior_per_tok_lst = game_segments[i].llm_prior_per_tok_segment[beg_index:end_index] + pad_cot_prefix_lst = game_segments[i].cot_prefix_segment[beg_index:end_index] + pad_llm_action_lst = game_segments[i].llm_action_segment[beg_index:end_index] + + pad_action_lst = game_segments[i].action_segment[:self.policy_config.num_unroll_steps + self.policy_config.td_steps] + pad_child_visits_lst = game_segments[i].child_visit_segment[:self.policy_config.num_unroll_steps + self.policy_config.td_steps] + + beg_index = 0 + end_index = beg_index + self.unroll_plus_td_steps - 1 + pad_reward_lst = game_segments[i].reward_segment[beg_index:end_index] + + if self.policy_config.use_ture_chance_label_in_chance_encoder: + chance_lst = game_segments[i].chance_segment[beg_index:end_index] + + beg_index = 0 + end_index = beg_index + self.unroll_plus_td_steps + pad_root_values_lst = game_segments[i].root_value_segment[beg_index:end_index] + + if self.policy_config.gumbel_algo: + pad_improved_policy_prob = game_segments[i].improved_policy_probs[beg_index:end_index] + + # Pad and finalize + if self.policy_config.gumbel_algo: + last_game_segments[i].pad_over( + pad_obs_lst, pad_reward_lst, pad_action_lst, pad_root_values_lst, pad_child_visits_lst, + next_segment_improved_policy=pad_improved_policy_prob, + next_segment_cot_prefix=pad_cot_prefix_lst, + next_segment_llm_action=pad_llm_action_lst + ) + else: + if self.policy_config.use_ture_chance_label_in_chance_encoder: + last_game_segments[i].pad_over( + pad_obs_lst, pad_reward_lst, pad_action_lst, pad_root_values_lst, pad_child_visits_lst, + next_chances=chance_lst, next_segment_raw_obs=pad_raw_obs_lst, + next_segment_history_obs=pad_history_obs_lst, next_segment_llm_prior_per_tok=pad_llm_prior_per_tok_lst, + next_segment_cot_prefix=pad_cot_prefix_lst, + next_segment_llm_action=pad_llm_action_lst + ) + else: + last_game_segments[i].pad_over( + pad_obs_lst, pad_reward_lst, pad_action_lst, pad_root_values_lst, pad_child_visits_lst, + next_segment_raw_obs=pad_raw_obs_lst, next_segment_history_obs=pad_history_obs_lst, + next_segment_llm_prior_per_tok=pad_llm_prior_per_tok_lst, + next_segment_cot_prefix=pad_cot_prefix_lst, + next_segment_llm_action=pad_llm_action_lst + ) + + last_game_segments[i].game_segment_to_array() + self.game_segment_pool.append((last_game_segments[i], last_game_priorities[i], done)) + + last_game_segments[i] = None + last_game_priorities[i] = None + + def _compute_priorities(self, game_segment: GameSegment) -> np.ndarray: + """Compute priorities for the game segment.""" + # Simple priority: uniform for now + return np.ones(len(game_segment.reward_segment)) + + def apply_temperature_scaling(self, prior: np.ndarray, return_logprobs: bool = False) -> np.ndarray: + """Apply temperature scaling to prior distribution.""" + if return_logprobs: + # Convert to log probabilities + log_probs = np.log(prior + 1e-10) + return log_probs + else: + return prior diff --git a/zoo/jericho/priorzero/priorzero_config.py b/zoo/jericho/priorzero/priorzero_config.py deleted file mode 100644 index 1614aaed4..000000000 --- a/zoo/jericho/priorzero/priorzero_config.py +++ /dev/null @@ -1,688 +0,0 @@ -# priorzero_config.py -""" -[PRIORZERO] PriorZero Configuration - -This module provides complete configuration for PriorZero algorithm. - -Key Features: -- Complete UniZero world model configuration -- LLM policy configuration (ORZ-style) -- Action space mapping for text environments -- Flexible switches to enable/disable components - -Author: PriorZero Team -Date: 2025-01-20 -""" - -import os -from typing import Dict, Tuple -from easydict import EasyDict - - -def get_jericho_action_mapping(env_id: str = 'zork1.z5') -> Tuple[Dict[str, int], Dict[int, str]]: - """ - Get action mapping for Jericho environments. - - In Jericho, the action space is typically defined by the game's valid actions. - For simplicity, we'll provide a basic mapping that can be extended. - - Args: - env_id: Jericho game ID - - Returns: - action_map: Mapping from action text to action index - action_inv_map: Mapping from action index to action text - """ - # Basic common actions for text adventure games - # These should ideally be loaded from the environment's action space - common_actions = [ - # Movement - "go north", "go south", "go east", "go west", - "go up", "go down", "go northeast", "go northwest", - "go southeast", "go southwest", - # Object interaction - "take all", "drop all", "inventory", "look", - "examine", "open", "close", "unlock", - # Common verbs - "read", "eat", "drink", "wear", "remove", - ] - - # Create mapping - action_map = {action.lower(): idx for idx, action in enumerate(common_actions)} - action_inv_map = {idx: action for action, idx in action_map.items()} - - return action_map, action_inv_map - - -def get_priorzero_config( - env_id: str = 'zork1.z5', - seed: int = 0, - exp_name: str = None, - enable_llm: bool = True, - enable_rft: bool = True, - debug_mode: bool = False, -) -> Tuple[EasyDict, EasyDict]: - """ - Generate complete PriorZero configuration. - - Args: - env_id: Jericho game ID - seed: Random seed - exp_name: Experiment name (auto-generated if None) - enable_llm: Whether to enable LLM policy (if False, degrades to pure UniZero) - enable_rft: Whether to enable RFT training (if False, only use SFT) - debug_mode: Whether to enable detailed debug logging (obs, action, LLM output, etc.) - - Returns: - main_config: Main configuration dictionary - create_config: Creation configuration for DI-engine components - """ - - # ============================================================================== - # 1. Basic Settings - # ============================================================================== - # Action space and max steps per environment (from jericho_unizero_config.py) - env_configurations = { - 'detective.z5': (12, 100), - 'omniquest.z5': (25, 100), - 'acorncourt.z5': (45, 50), - 'zork1.z5': (55, 500), - } - action_space_size, max_steps = env_configurations.get(env_id, (20, 100)) - - # World model encoder (for processing text observations) - wm_encoder_option = 'legacy' # Options: 'legacy', 'clip', 'custom' - wm_model_name = 'BAAI/bge-base-en-v1.5' # Sentence transformer for text encoding - - # LLM policy model - # llm_model_name = "Qwen/Qwen2.5-1.5B-Instruct" # Smaller model for faster iteration - llm_model_name = "Qwen/Qwen2.5-0.5B-Instruct" # Smaller model for faster iteration - - # Get action mappings - action_map, action_inv_map = get_jericho_action_mapping(env_id) - - # Convert action_inv_map to use string keys for EasyDict compatibility - action_inv_map_str = {str(k): v for k, v in action_inv_map.items()} - - # ============================================================================== - # 2. Environment Configuration - # ============================================================================== - env_config = dict( - # Stop conditions - stop_value=int(1e6), - max_steps=max_steps, - - # Observation and action space - observation_shape=512, # BGE embedding dimension - action_space_size=action_space_size, - - # [FIX] Jericho environment expects these at top level - env_id=env_id, - game_path=f"/mnt/nfs/zhangjinouwen/puyuan/LightZero/zoo/jericho/envs/z-machine-games-master/jericho-game-suite/{env_id}", - tokenizer_path=wm_model_name, - env_type="jericho", - max_action_num=action_space_size, - max_seq_len=512, - save_replay=False, - save_replay_path="", - collect_policy_mode="default", - - # Parallelization - collector_env_num=4, - evaluator_env_num=2, - n_evaluator_episode=2, - - # Environment manager - manager=dict( - shared_memory=False, - reset_timeout=60, # Increased timeout for text env initialization - ), - ) - - # ============================================================================== - # 3. UniZero World Model Configuration - # ============================================================================== - world_model_config = dict( - # [CRITICAL] DI-engine requires 'type' field to identify model class - type='UniZeroModel', - - # [FIX] EasyDict.pop() doesn't handle default values properly, must include import_names - import_names=[], # Empty list since UniZeroModel is already registered - - # Model type - model_type='mlp', # For vector observations (text embeddings) - continuous_action_space=False, - - # Observation and action - observation_shape=512, - action_space_size=action_space_size, - - # [FIX] Encoder settings must be at top level for UniZeroModel.__init__ - encoder_option=wm_encoder_option, - encoder_url=wm_model_name, - - # World model architecture - world_model_cfg=dict( - # Obs type - obs_type="text", # Important: text-based observations - - # Environment settings - env_num=max(4, 2), # max(collector_env_num, evaluator_env_num), will be updated in quick_test - action_space_size=action_space_size, - - # Transformer settings - # num_layers=4, # Reduced for faster training - num_layers=2, # Reduced for faster training # TODO - num_heads=8, - embed_dim=512, - - # Context and unroll - # Note: Each timestep contains 2 tokens: observation and action - num_unroll_steps=10, # Number of steps to unroll in training - infer_context_length=4, # Inference context length - tokens_per_block=2, # obs + action - max_blocks=10, # num_unroll_steps (default) - max_tokens=2 * 10, # 2 * num_unroll_steps - context_length=2 * 4, # 2 * infer_context_length - - # Regularization - embed_pdrop=0.1, - resid_pdrop=0.1, - attn_pdrop=0.1, - - # Loss weights - latent_recon_loss_weight=0.0, # Latent reconstruction loss - perceptual_loss_weight=0.0, - policy_entropy_weight=0.0, # Entropy regularization - - # Normalization - final_norm_option_in_head="LayerNorm", - final_norm_option_in_encoder="LayerNorm", - predict_latent_loss_type='mse', # or 'group_kl' with SimNorm - - # Device - device="cuda", - - # Advanced settings - gru_gating=False, - attention='causal', - support_size=101, # For distributional RL - - # Analysis flags - analysis_sim_norm=False, - analysis_dormant_ratio_weight_rank=False, - # use_priority=False, - use_priority=True, - - # Position encoding - rotary_emb=False, # Whether to use RoPE - rope_theta=10000, - max_seq_len=8192, - - # LoRA (optional, for world model) - lora_r=0, # Set > 0 to enable LoRA - - # Other - decode_loss_mode=None, # 'after_backbone', 'before_backbone', or None - gamma=1.0, # Discount factor - dormant_threshold=0.025, - - task_embed_option=None, - use_task_embed=False, - use_normal_head=True, - use_softmoe_head=False, - use_moe_head=False, - num_experts_in_moe_head=4, - moe_in_transformer=False, - multiplication_moe_in_transformer=False, - n_shared_experts=1, - num_experts_per_tok=1, - num_experts_of_moe_in_transformer=8, - # game_segment_length=200, - game_segment_length=50, - ), - - # Distributional RL - categorical_distribution=True, - reward_support_range=(-50., 51., 1.), # (min, max, step) for reward support - value_support_range=(-50., 51., 1.), # (min, max, step) for value support - - # Self-supervised learning - self_supervised_learning_loss=True, - - # Model architecture details - frame_stack_num=1, - bias=True, - res_connection_in_dynamics=True, - norm_type='LN', # LayerNorm for text - ) - - # ============================================================================== - # 4. LLM Policy Configuration (ORZ-style) - # ============================================================================== - llm_policy_config = dict( - # Model path - pretrain_llm_path=llm_model_name, - - # LoRA for parameter-efficient fine-tuning - use_lora=False, # Set to True to enable LoRA - lora_r=8, - lora_alpha=16, - lora_dropout=0.05, - - # Training - llm_learning_rate=1e-6, - llm_weight_decay=0.01, - llm_loss_weight=0.5, # Weight of SFT loss in total loss - rft_loss_weight=0.3, # Weight of RFT loss in total loss - - # [PRIORZERO-OOM-FIX] Gradient accumulation for memory efficiency - # Process LLM training in smaller micro-batches to avoid OOM - llm_micro_batch_size=4, # Small batch size per forward pass (reduce if still OOM) - llm_gradient_accumulation_steps=8, # Accumulate gradients over 8 steps (effective batch = 4*8=32) - # Note: Effective batch size = llm_micro_batch_size * llm_gradient_accumulation_steps - - # Generation - prompt_max_len=2048, - generate_max_len=256, # Max tokens for LLM output - - # Prompting strategy - history_length=5, # Number of recent (obs, action, reward) tuples to include - use_cot=True, # Whether to use Chain-of-Thought prompting - - # Training strategy - sft_target='mcts_policy', # 'mcts_policy' or 'oracle_policy' - enable_rft=enable_rft, # Whether to enable RFT with env rewards - # enable_rft=False, # Whether to enable RFT with env rewards # TODO - - # vLLM settings - vllm_tensor_parallel_size=1, - gpu_memory_utilization=0.3, # Adjust based on your GPU memory - ) - - # ============================================================================== - # 5. Policy Configuration (Combines World Model + LLM) - # ============================================================================== - policy_config = dict( - learn=dict( - learner=dict( - hook=dict( - save_ckpt_after_iter=1000000, # To save memory, set a large value. If intermediate checkpoints are needed, reduce this value. - ), - ), - ), - type='priorzero', - - # Environment settings (must match env config) - collector_env_num=env_config['collector_env_num'], - evaluator_env_num=env_config['evaluator_env_num'], - - # Model config (world model) - model=world_model_config, - - # [PRIORZERO-NEW] LLM policy config - llm_policy_cfg=llm_policy_config, - - # [PRIORZERO-NEW] Action mappings (use original dict, not EasyDict) - # These will be set directly on policy instance, not through EasyDict - _action_map=action_map, # Prefix with _ to avoid EasyDict conversion - _action_inv_map=action_inv_map, - - # ============================================================================== - # [ASYNC-NEW] Async Training Configuration - # ============================================================================== - # off_policy_degree controls the degree of asynchrony between collect and train: - # - 0: Fully synchronous (serial) mode - collect -> train -> eval - # - 1-10: Low async - train can lag behind collect by a few batches - # - 10-50: Medium async - train can lag more, higher throughput - # - >50: High async - maximum throughput, highest off-policy bias - # - # Special value -1: Auto-tune based on buffer size and batch size - off_policy_degree=0, # Default to synchronous mode for stability - # off_policy_degree=5, - - # Whether to enable async evaluation (runs eval in background) - enable_async_eval=False, - - # MCTS settings - num_simulations=25, - collect_num_simulations=25, - eval_num_simulations=25, - - # MCTS exploration - root_dirichlet_alpha=0.3, - root_noise_weight=0.25, - - # MCTS variants (set one to True to use that variant) - sampled_algo=False, # Sampled MuZero - gumbel_algo=False, # Gumbel MuZero - mcts_ctree=True, # Use C++ MCTS (faster) - - # Training settings - batch_size=32, - learning_rate=3e-4, # World model learning rate - weight_decay=1e-4, - optim_type='AdamW', - grad_clip_value=10.0, - - # Loss components - value_loss_weight=1.0, - policy_loss_weight=1.0, - reward_loss_weight=1.0, - - # Adaptive entropy weight (for exploration) - use_adaptive_entropy_weight=True, - adaptive_entropy_alpha_lr=1e-4, - - # Encoder gradient clipping with annealing - use_encoder_clip_annealing=True, - encoder_clip_anneal_type='cosine', - encoder_clip_start_value=30.0, - encoder_clip_end_value=10.0, - encoder_clip_anneal_steps=100000, - - # Training schedule - num_unroll_steps=10, - td_steps=5, - train_start_after_envsteps=0, - # train_start_after_envsteps=1000, - update_per_collect=None, # Will be set automatically - replay_ratio=0.25, - - # Replay buffer - # replay_buffer_size=int(1e4), - replay_buffer_size=int(1e5), - use_priority=True, # Prioritized experience replay - priority_prob_alpha=0.6, - priority_prob_beta=0.4, - - # Evaluation - eval_freq=500, - - # Game segments - # game_segment_length=200, - game_segment_length=50, - num_segments=env_config['collector_env_num'], # Must equal collector_env_num - - # Misc - ignore_done=False, - collect_with_pure_policy=False, - monitor_extra_statistics=True, - - # Device - cuda=True, - device='cuda', - multi_gpu=False, - - # Environment type - env_type='not_board_games', - action_type='varied_action_space', # Jericho has varied action space per state - battle_mode='play_with_bot_mode', - - # Data processing - transform2string=False, - gray_scale=False, - use_augmentation=False, - - # Advanced - use_rnd_model=False, # Random Network Distillation for exploration - analysis_sim_norm=False, - sample_type='transition', - - # ============================================================================== - # [ALIGN WITH UNIZERO] Reanalyze Configuration (atari_unizero_segment_config.py line 201-206) - # ============================================================================== - # Defines the frequency of reanalysis. E.g., 1 means reanalyze once per epoch, - # 2 means reanalyze once every two epochs, 1/50 means reanalyze once every 50 epochs. - buffer_reanalyze_freq=1/5000000000, # Effectively disabled for Jericho (set very low) - # Each reanalyze process will reanalyze sequences - # ( transitions per sequence) - reanalyze_batch_size=160, - # The partition of reanalyze. E.g., 1 means reanalyze_batch samples from the whole buffer, - # 0.5 means samples from the first half of the buffer. - reanalyze_partition=0.75, - # Reanalyze ratio (used in some algorithms, kept for compatibility) - reanalyze_ratio=0.0, - ) - - # ============================================================================== - # 6. Replay Buffer Configuration - # ============================================================================== - replay_buffer_config = dict( - type='game', - replay_buffer_size=policy_config['replay_buffer_size'], - batch_size=policy_config['batch_size'], - ) - - # ============================================================================== - # 6.5 Remove problematic nested dicts before EasyDict conversion - # ============================================================================== - # Store action mappings separately to avoid EasyDict issues with integer keys - _temp_action_map = action_map - _temp_action_inv_map = action_inv_map - - # ============================================================================== - # 7. Main Configuration Assembly - # ============================================================================== - priorzero_config = dict( - env=env_config, - policy=policy_config, - replay_buffer=replay_buffer_config, - - # Experiment settings - exp_name=exp_name or f"priorzero_{env_id}_seed{seed}", - seed=seed, - - # Debug settings - debug_mode=debug_mode, - ) - - # ============================================================================== - # 8. Create Configuration (for DI-engine component creation) - # ============================================================================== - create_config = dict( - env=dict( - type="jericho", - import_names=["zoo.jericho.envs.jericho_env"], - ), - env_manager=dict( - type="base" # [FIX] Use 'base' for jericho to avoid daemon process issues - ), - policy=dict( - type="priorzero", - import_names=["zoo.jericho.priorzero.priorzero_policy"], - ), - collector=dict( - type="priorzero_segment", - import_names=["zoo.jericho.priorzero.priorzero_collector"], - ), - evaluator=dict( - type="priorzero", - import_names=["zoo.jericho.priorzero.priorzero_evaluator"], - ), - replay_buffer=dict( - type='game_buffer_muzero', - import_names=['lzero.mcts.buffer.game_buffer_muzero'], - ), - ) - - # ============================================================================== - # 9. Convert to EasyDict for convenient access - # ============================================================================== - # IMPORTANT: Remove _action_map and _action_inv_map from policy_config before EasyDict - # to avoid integer key issues - policy_config_copy = {k: v for k, v in policy_config.items() if not k.startswith('_')} - priorzero_config['policy'] = policy_config_copy - - main_config = EasyDict(priorzero_config) - create_config = EasyDict(create_config) - - # Set experiment path - main_config.exp_name = f"data_priorzero/{main_config.exp_name}" - - # [IMPORTANT] Set action mappings as regular attributes (not through EasyDict) - # Use object.__setattr__ to bypass EasyDict's __setattr__ which tries to convert dicts - object.__setattr__(main_config.policy, 'action_map', _temp_action_map) - object.__setattr__(main_config.policy, 'action_inv_map', _temp_action_inv_map) - - return main_config, create_config - - -def get_priorzero_config_for_quick_test(env_id: str = 'zork1.z5', seed: int = 0, debug_mode: bool = False): - """ - Get a lightweight configuration for quick testing (reduced resources). - - This is useful for: - - Debugging - - CI/CD pipelines - - Local development without powerful GPUs - - IMPORTANT: All sequence-length related parameters must be consistent: - - num_unroll_steps: Number of timesteps in training unroll - - max_blocks: Should equal num_unroll_steps - - max_tokens: Should equal num_unroll_steps * tokens_per_block (= num_unroll_steps * 2) - - infer_context_length: Context length for inference - - context_length: Should equal infer_context_length * tokens_per_block (= infer_context_length * 2) - """ - main_config, create_config = get_priorzero_config(env_id, seed, debug_mode=debug_mode) - - # ============================================================================== - # [CRITICAL FIX] Define num_unroll_steps FIRST to ensure consistency - # ============================================================================== - quick_test_num_unroll_steps = 10 # Core parameter that determines sequence length - quick_test_infer_context_length = 4 # Inference context length - tokens_per_block = 2 # obs + action (fixed in UniZero architecture) - - # Reduce computational requirements - main_config.env.collector_env_num = 2 - main_config.env.evaluator_env_num = 1 - main_config.env.n_evaluator_episode = 1 - - # ============================================================================== - # Policy-level configurations - # ============================================================================== - main_config.policy.num_simulations = 5 - # main_config.policy.batch_size = 20 - main_config.policy.batch_size = 2 - main_config.policy.game_segment_length = 20 # Can be larger than num_unroll_steps - main_config.policy.num_segments = 2 # Must equal collector_env_num - main_config.policy.replay_buffer_size = 1000 - - # [CRITICAL] Set policy-level num_unroll_steps to match world model - main_config.policy.num_unroll_steps = quick_test_num_unroll_steps - - # ============================================================================== - # World model configurations - ALL must be consistent with num_unroll_steps - # ============================================================================== - main_config.policy.model.world_model_cfg.num_layers = 1 - main_config.policy.model.world_model_cfg.num_heads = 2 - - # Update env_num to match the reduced collector/evaluator counts - main_config.policy.model.world_model_cfg.env_num = max( - main_config.env.collector_env_num, - main_config.env.evaluator_env_num - ) - - # [CRITICAL] Sequence length parameters - must all be consistent - main_config.policy.model.world_model_cfg.num_unroll_steps = quick_test_num_unroll_steps - main_config.policy.model.world_model_cfg.max_blocks = quick_test_num_unroll_steps - main_config.policy.model.world_model_cfg.max_tokens = quick_test_num_unroll_steps * tokens_per_block # 3 * 2 = 6 - - main_config.policy.model.world_model_cfg.infer_context_length = quick_test_infer_context_length - main_config.policy.model.world_model_cfg.context_length = quick_test_infer_context_length * tokens_per_block # 2 * 2 = 4 - - # Verify tokens_per_block is set correctly (should already be 2 from base config) - main_config.policy.model.world_model_cfg.tokens_per_block = tokens_per_block - - # ============================================================================== - # LLM policy configurations - # ============================================================================== - main_config.policy.llm_policy_cfg.prompt_max_len = 1024 - main_config.policy.llm_policy_cfg.generate_max_len = 128 - main_config.policy.llm_policy_cfg.history_length = 3 - # [PRIORZERO-OOM-FIX] Reduce micro-batch size for quick test to avoid OOM - main_config.policy.llm_policy_cfg.llm_micro_batch_size = 2 - main_config.policy.llm_policy_cfg.llm_gradient_accumulation_steps = 4 - - main_config.exp_name = f"{main_config.exp_name}_debug" - - return main_config, create_config - - -# ============================================================================== -# Preset Configurations for Different Scenarios -# ============================================================================== - -def get_config_pure_unizero(env_id: str = 'zork1.z5', seed: int = 0): - """Get config for pure UniZero (without LLM).""" - main_config, create_config = get_priorzero_config( - env_id=env_id, - seed=seed, - enable_llm=False, - ) - main_config.exp_name = f"pure_unizero_{env_id}_seed{seed}" - main_config.policy.llm_policy_cfg.llm_loss_weight = 0.0 - main_config.policy.llm_policy_cfg.rft_loss_weight = 0.0 - return main_config, create_config - - -def get_config_llm_only_sft(env_id: str = 'zork1.z5', seed: int = 0): - """Get config for LLM with only SFT (no RFT).""" - main_config, create_config = get_priorzero_config( - env_id=env_id, - seed=seed, - enable_rft=False, - ) - main_config.exp_name = f"priorzero_sft_only_{env_id}_seed{seed}" - return main_config, create_config - - -def get_config_with_lora(env_id: str = 'zork1.z5', seed: int = 0): - """Get config with LoRA enabled for LLM (memory efficient).""" - main_config, create_config = get_priorzero_config(env_id=env_id, seed=seed) - main_config.policy.llm_policy_cfg.use_lora = True - main_config.exp_name = f"priorzero_lora_{env_id}_seed{seed}" - return main_config, create_config - - -# ============================================================================== -# Example Usage -# ============================================================================== - -if __name__ == "__main__": - # Test configuration generation - print("="*80) - print("Testing PriorZero Configuration Generation") - print("="*80) - - # 1. Standard config - print("\n1. Standard PriorZero Config:") - main_cfg, create_cfg = get_priorzero_config(env_id='zork1.z5', seed=0) - print(f" Exp name: {main_cfg.exp_name}") - print(f" Action space size: {main_cfg.policy.model.action_space_size}") - print(f" LLM model: {main_cfg.policy.llm_policy_cfg.pretrain_llm_path}") - print(f" World model layers: {main_cfg.policy.model.world_model_cfg.num_layers}") - print(f" Num action mappings: {len(main_cfg.policy.action_map)}") - - # 2. Quick test config - print("\n2. Quick Test Config:") - test_cfg, _ = get_priorzero_config_for_quick_test() - print(f" Batch size: {test_cfg.policy.batch_size}") - print(f" Num simulations: {test_cfg.policy.num_simulations}") - print(f" Collector envs: {test_cfg.env.collector_env_num}") - - # 3. Pure UniZero config - print("\n3. Pure UniZero Config:") - unizero_cfg, _ = get_config_pure_unizero() - print(f" LLM loss weight: {unizero_cfg.policy.llm_policy_cfg.llm_loss_weight}") - print(f" RFT enabled: {unizero_cfg.policy.llm_policy_cfg.enable_rft}") - - # 4. Config with LoRA - print("\n4. Config with LoRA:") - lora_cfg, _ = get_config_with_lora() - print(f" Use LoRA: {lora_cfg.policy.llm_policy_cfg.use_lora}") - print(f" LoRA rank: {lora_cfg.policy.llm_policy_cfg.lora_r}") - - print("\n" + "="*80) - print("✓ All configurations generated successfully!") - print("="*80) diff --git a/zoo/jericho/priorzero/priorzero_datafactory_unified.py b/zoo/jericho/priorzero/priorzero_datafactory_unified.py new file mode 100644 index 000000000..a3bc8993d --- /dev/null +++ b/zoo/jericho/priorzero/priorzero_datafactory_unified.py @@ -0,0 +1,681 @@ +""" +Unified DataProcessor supporting both text (LLM) and image (VL) inputs + +This processor can handle: +- Text observations with LLM (original functionality) +- Image observations with VL (new functionality) +""" +from __future__ import annotations +from dataclasses import dataclass +from typing import Any, Dict, List, Optional, Tuple, Union +import re +import random +import torch +import torch.distributed as dist +from vllm import SamplingParams +from ding.utils import build_logger +import numpy as np +from PIL import Image + + +class UnifiedDataProcessor: + """ + Unified DataProcessor supporting both text and image inputs. + + For text input: Uses LLM (vLLM engine) + For image input: Uses VL engine + """ + + def __init__( + self, + rank: int, + world_size: int, + vllm_engine, # Can be vLLM or VL engine + strategy, + model_path: str, + exp_name: Optional[str] = None, + instance_name: str = "unified_output", + obs_type: str = 'text', # NEW: 'text' or 'image' + ): + """ + Initialize Unified DataProcessor. + + Args: + rank: Process rank + world_size: World size + vllm_engine: vLLM or VL engine + strategy: Training strategy + model_path: Model path + exp_name: Experiment name + instance_name: Instance name for logging + obs_type: Observation type ('text' or 'image') + """ + self.vllm_engine = vllm_engine + self.strategy = strategy + self.args = getattr(strategy, "args", None) + self.obs_type = obs_type # NEW + + # Load tokenizer + from transformers import AutoTokenizer + self.tokenizer = AutoTokenizer.from_pretrained( + model_path, trust_remote_code=True, padding_side="left" + ) + if self.tokenizer.pad_token is None: + self.tokenizer.pad_token = self.tokenizer.eos_token + + # Configuration + self.use_cot = getattr(self.args, 'use_cot', True) + self.prompt_max_len = getattr(self.args, 'prompt_max_len', 8192) + self.generate_max_len = getattr(self.args, 'generate_max_len', 512) + self.temperature = getattr(self.args, 'temperature', 1.0) + self.top_p = getattr(self.args, 'top_p', 1.0) + self.vllm_enable_sleep = getattr(self.args, 'vllm_enable_sleep', True) + self.reduction = getattr(self.args, 'reduction', 'mean') + self.rank = rank + self.world_size = world_size + self.output_step = 0 + self.llm_prior_with_cot = False + + # Statistics + self.episode_output = [] + self.value_running_mean = 0.0 + self.value_running_std = 1.0 + self.value_count = 0 + self.running_momentum = 0.99 + + # Logger + if self.rank == 0: + self._logger, _ = build_logger( + path=f'./{exp_name}/log/{instance_name}', + name=instance_name, + need_tb=False + ) + self._logger.info(f"✓ UnifiedDataProcessor initialized") + self._logger.info(f" - Observation type: {obs_type}") + self._logger.info(f" - Use CoT: {self.use_cot}") + + # Value normalizer + if hasattr(self.args, 'value_norm_cfg') and self.args.value_norm_cfg.enable_stability_optimizer: + from models.stability_optimizer import AdaptiveValueNormalizer + self.value_normalizer = AdaptiveValueNormalizer( + init_momentum=self.args.value_norm_cfg.value_norm_init_momentum, + final_momentum=self.args.value_norm_cfg.value_norm_final_momentum, + warmup_steps=self.args.value_norm_cfg.value_norm_warmup_steps, + clip_method=self.args.value_norm_cfg.value_norm_clip_method, + clip_percentile=self.args.value_norm_cfg.value_norm_clip_percentile, + min_std=1e-6, + history_size=self.args.value_norm_cfg.value_norm_history_size, + ) + else: + self.value_normalizer = None + + # ========================================================================= + # Text Input Methods (Original LLM functionality) + # ========================================================================= + + def get_system_prompt_text(self) -> str: + """System prompt for text-based games (LLM).""" + parts = [ + "You are an expert player in a text-based adventure game.", + "Your goal is to maximize the score by choosing the optimal next action.", + "Please analyze the game history and current observation to decide the single best next action.", + "OUTPUT FORMAT:", + ] + + if self.use_cot: + parts.append( + "You MUST produce exactly TWO parts in the following order:\n" + "1. Reasoning: Analyze the current situation, available actions, constraints, and uncertainties.\n" + "2. Action: The final chosen action.\n" + "Strict Format Example:\n" + "Reasoning: \n" + "Action: " + ) + else: + parts.append( + "Output exactly one line starting with 'Action:'.\n" + "Example:\n" + "Action: " + ) + return "\n".join(parts) + + def get_user_prompt_text( + self, + history: Optional[List[Tuple[str, str, float]]] = None, + current_obs: Optional[str] = None, + valid_actions: Optional[List[str]] = None + ) -> str: + """User prompt for text-based games (LLM).""" + prompt_parts = [] + + if history and len(history) > 0: + prompt_parts.append("=== GAME HISTORY ===") + for i, (obs, action, reward) in enumerate(history, start=1): + prompt_parts.append(f"Step {i}:") + prompt_parts.append(f"Observation: {obs.strip()}") + prompt_parts.append(f"Action: {action.strip()}") + prompt_parts.append(f"Reward: {reward}") + prompt_parts.append("") + + prompt_parts.append("=== CURRENT OBSERVATION ===") + prompt_parts.append(current_obs.strip()) + + if valid_actions: + prompt_parts.append("\n=== VALID ACTIONS ===") + for i, action in enumerate(valid_actions, start=1): + prompt_parts.append(f"{i}. {action}") + + prompt_parts.append("\n=== INSTRUCTION ===") + prompt_parts.append("Choose the best action from the valid actions above.") + + return "\n".join(prompt_parts) + + # ========================================================================= + # Image Input Methods (NEW VL functionality) + # ========================================================================= + + def get_system_prompt_image(self) -> str: + """System prompt for image-based games (VL).""" + parts = [ + "You are an expert Atari game player.", + "Your goal is to maximize the score by choosing the optimal next action based on the game screen.", + "Analyze the current game state shown in the image and decide the best action.", + "OUTPUT FORMAT:", + ] + + if self.use_cot: + parts.append( + "You MUST produce exactly TWO parts:\n" + "1. Reasoning: Analyze the game state (positions, velocities, score, etc.)\n" + "2. Action: The final chosen action.\n" + "Format:\n" + "Reasoning: \n" + "Action: " + ) + else: + parts.append( + "Output exactly one line starting with 'Action:'.\n" + "Example:\n" + "Action: " + ) + return "\n".join(parts) + + def get_user_prompt_image( + self, + history: Optional[List[Tuple[Any, str, float]]] = None, + valid_actions: Optional[List[str]] = None, + game_context: Optional[str] = None + ) -> str: + """User prompt for image-based games (VL).""" + prompt_parts = [] + + if game_context: + prompt_parts.append(f"=== GAME CONTEXT ===") + prompt_parts.append(game_context) + prompt_parts.append("") + + if history and len(history) > 0: + prompt_parts.append("=== RECENT HISTORY ===") + for i, (_, action, reward) in enumerate(history[-3:], start=1): # Last 3 steps + prompt_parts.append(f"Step {i}: Action={action}, Reward={reward}") + prompt_parts.append("") + + prompt_parts.append("=== CURRENT GAME SCREEN ===") + prompt_parts.append("(See the image above)") + + if valid_actions: + prompt_parts.append("\n=== VALID ACTIONS ===") + for i, action in enumerate(valid_actions, start=1): + prompt_parts.append(f"{i}. {action}") + + prompt_parts.append("\n=== INSTRUCTION ===") + prompt_parts.append("Based on the current game screen, choose the best action from the valid actions above.") + + return "\n".join(prompt_parts) + + # ========================================================================= + # Unified Interface + # ========================================================================= + + def get_action_prior_single( + self, + observation: Union[str, np.ndarray, Image.Image], + action_candidates: List[str], + history: Optional[List] = None, + temperature: float = 1.0, + use_cot: Optional[bool] = None, + ) -> Dict[str, Any]: + """ + Get action prior for a single observation (unified interface). + + Args: + observation: Text string or image array/PIL Image + action_candidates: List of valid actions + history: Optional history + temperature: Sampling temperature + use_cot: Whether to use CoT (overrides self.use_cot) + + Returns: + Dictionary with action_probs, action_logits, raw_output + """ + if use_cot is None: + use_cot = self.use_cot + + if self.obs_type == 'text': + return self._get_action_prior_text( + text_obs=observation, + action_candidates=action_candidates, + history=history, + temperature=temperature, + use_cot=use_cot, + ) + else: # image + return self._get_action_prior_image( + image_obs=observation, + action_candidates=action_candidates, + history=history, + temperature=temperature, + use_cot=use_cot, + ) + + def _get_action_prior_text( + self, + text_obs: str, + action_candidates: List[str], + history: Optional[List] = None, + temperature: float = 1.0, + use_cot: bool = True, + ) -> Dict[str, Any]: + """Get action prior for text observation using LLM.""" + # Build prompt + system_prompt = self.get_system_prompt_text() + user_prompt = self.get_user_prompt_text( + history=history, + current_obs=text_obs, + valid_actions=action_candidates + ) + + # Build chat messages + messages = [ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": user_prompt} + ] + + # Convert to text + prompt_text = self.tokenizer.apply_chat_template( + messages, + tokenize=False, + add_generation_prompt=True + ) + + # Generate with vLLM + sampling_params = SamplingParams( + temperature=temperature, + top_p=self.top_p, + max_tokens=self.generate_max_len, + ) + + outputs = self.vllm_engine.generate([prompt_text], sampling_params) + raw_output = outputs[0].outputs[0].text + + # Parse output to get action probabilities + action_probs = self._parse_llm_output_to_probs(raw_output, action_candidates) + action_logits = np.log(action_probs + 1e-10) + + return { + 'action_probs': action_probs, + 'action_logits': action_logits, + 'raw_output': raw_output, + } + + def _get_action_prior_image( + self, + image_obs: Union[np.ndarray, Image.Image], + action_candidates: List[str], + history: Optional[List] = None, + temperature: float = 1.0, + use_cot: bool = True, + ) -> Dict[str, Any]: + """Get action prior for image observation using VL.""" + # Convert to PIL Image if needed + if isinstance(image_obs, np.ndarray): + if image_obs.dtype != np.uint8: + image_obs = (image_obs * 255).astype(np.uint8) + # Handle different formats + if image_obs.shape[0] == 3: # (C, H, W) -> (H, W, C) + image_obs = np.transpose(image_obs, (1, 2, 0)) + image = Image.fromarray(image_obs) + else: + image = image_obs + + # Build prompt + system_prompt = self.get_system_prompt_image() + user_prompt = self.get_user_prompt_image( + history=history, + valid_actions=action_candidates, + game_context="Atari game" + ) + + # Combine prompts + full_prompt = f"{system_prompt}\n\n{user_prompt}" + + # Generate with VL + raw_output = self.vllm_engine.generate( + image=image, + prompt=full_prompt, + temperature=temperature, + max_new_tokens=self.generate_max_len, + ) + + # Parse output to get action probabilities + action_probs = self._parse_vl_output_to_probs(raw_output, action_candidates) + action_logits = np.log(action_probs + 1e-10) + + return { + 'action_probs': action_probs, + 'action_logits': action_logits, + 'raw_output': raw_output, + } + + def _parse_llm_output_to_probs(self, raw_output: str, action_candidates: List[str]) -> np.ndarray: + """Parse LLM output to action probabilities.""" + # Extract action from output + action_match = re.search(r'Action:\s*(.+)', raw_output, re.IGNORECASE) + if action_match: + chosen_action = action_match.group(1).strip() + + # Find matching action + for i, action in enumerate(action_candidates): + if action.lower() in chosen_action.lower() or chosen_action.lower() in action.lower(): + # High probability for chosen action + probs = np.ones(len(action_candidates)) * 0.01 + probs[i] = 0.9 + probs = probs / probs.sum() + return probs + + # Fallback: uniform distribution + return np.ones(len(action_candidates)) / len(action_candidates) + + def _parse_vl_output_to_probs(self, raw_output: str, action_candidates: List[str]) -> np.ndarray: + """Parse VL output to action probabilities.""" + # Similar to LLM parsing + return self._parse_llm_output_to_probs(raw_output, action_candidates) + + def get_llm_prior( + self, + states: List[Union[str, np.ndarray, Image.Image]], + valid_actions_list: List[List[str]], + histories: Optional[List[List]] = None, + return_cot: bool = False + ) -> Tuple[List[np.ndarray], List[np.ndarray], List[Any]]: + """ + Batch get LLM/VL priors (for backward compatibility). + + Args: + states: List of observations (text or images) + valid_actions_list: List of valid action lists + histories: List of histories + return_cot: Whether to return CoT prefixes + + Returns: + Tuple of (prior_per_seq, prior_per_tok, cot_prefixes) + """ + if histories is None: + histories = [None] * len(states) + + prior_per_seq = [] + prior_per_tok = [] + cot_prefixes = [] + + for obs, actions, hist in zip(states, valid_actions_list, histories): + result = self.get_action_prior_single( + observation=obs, + action_candidates=actions, + history=hist, + temperature=self.temperature, + ) + + prior_per_seq.append(result['action_probs']) + prior_per_tok.append(result['action_logits']) + if return_cot: + cot_prefixes.append(result['raw_output']) + + if return_cot: + return prior_per_seq, prior_per_tok, cot_prefixes + else: + return prior_per_seq, prior_per_tok, [None] * len(states) + + def make_llm_train_samples(self, priorzero_batch, ddp: bool = True, max_samples: int = None, prior_generator=None): + """ + Make training samples from PriorZero batch. + + For text mode: tokenizes prompts + actions → tensors expected by BatchPPOTrainer. + For image mode: builds VL chat context, tokenizes → same tensor format. + + Returns: + Tuple of (flag, train_samples) where train_samples is a 6-tuple: + (input_ids, attention_mask, action_mask, advantage, rollout_logprob, log_status) + """ + if self.obs_type == 'image': + return self._make_vl_train_samples(priorzero_batch, ddp=ddp, max_samples=max_samples, prior_generator=prior_generator) + else: + # Original LLM training samples (text input) + # Keep existing implementation + pass + + def _make_vl_train_samples(self, priorzero_batch, ddp: bool = True, max_samples: int = None, prior_generator=None): + """ + Build VL training samples in the same tensor format as the LLM path. + + The 8-element priorzero_batch from fetch_latest_batch: + [raw_obs_list, history_obs_list, llm_prior_per_tok_list, + batch_target_values, batch_pred_values, cot_prefix_list, llm_action_list, action_list] + + Returns: + (flag, (input_ids, attention_mask, action_mask, advantage, rollout_logprob, log_status)) + """ + import logging + import random + import traceback + _logger = logging.getLogger(__name__) + + try: + raw_obs_list, history_obs_list, llm_prior_per_tok_list, \ + target_values, pred_values, cot_prefix_list, llm_action_list, action_list = priorzero_batch + + if len(raw_obs_list) == 0: + return (False, []) + + B = len(raw_obs_list) + T = len(raw_obs_list[0]) if B > 0 else 0 + + # ---- Step 1: build flat sample list ---- + samples = [] + for b in range(B): + for t in range(T - 1): + action_name = llm_action_list[b][t + 1] + if action_name is None: + continue + + # history at time t = history after executing action t (before action t+1) + history = history_obs_list[b][t] if t < len(history_obs_list[b]) else [] + cot_prefix = cot_prefix_list[b][t + 1] if (cot_prefix_list is not None and t + 1 < len(cot_prefix_list[b])) else None + + # VL prior stored as action-level log-prob array + action_logprobs = llm_prior_per_tok_list[b][t + 1] if ( + llm_prior_per_tok_list is not None and t + 1 < len(llm_prior_per_tok_list[b]) + ) else None + + # MCTS-selected action index (integer) for correct rollout log-prob lookup + mcts_action_idx = int(action_list[b][t + 1]) if ( + action_list is not None and b < len(action_list) and t + 1 < len(action_list[b]) + ) else None + + tv = float(target_values[b][t]) if target_values is not None and b < len(target_values) and t < len(target_values[b]) else 0.0 + pv = float(pred_values[b][t]) if pred_values is not None and b < len(pred_values) and t < len(pred_values[b]) else 0.0 + + samples.append({ + 'history': history, + 'action_name': action_name, + 'cot_prefix': cot_prefix, + 'action_logprobs': action_logprobs, # np.ndarray or None + 'mcts_action_idx': mcts_action_idx, # int or None + 'target_value': tv, + 'pred_value': pv, + }) + + if len(samples) == 0: + return (False, []) + + random.Random(0).shuffle(samples) + + if max_samples is not None and len(samples) > max_samples: + samples = samples[:max_samples] + + if ddp: + real_samples = samples + else: + per_rank = len(samples) // self.world_size + start = self.rank * per_rank + end = (self.rank + 1) * per_rank if self.rank != self.world_size - 1 else len(samples) + real_samples = samples[start:end] + + if len(real_samples) == 0: + return (False, []) + + # ---- Step 2: build target text for each sample ---- + if self.use_cot: + targets_only = [] + for s in real_samples: + cot = s['cot_prefix'] or "" + if cot: + targets_only.append(cot.strip() + "\nAction: " + s['action_name'] + self.tokenizer.eos_token) + else: + targets_only.append("Action: " + s['action_name'] + self.tokenizer.eos_token) + else: + targets_only = ["Action: " + s['action_name'] + self.tokenizer.eos_token for s in real_samples] + + # ---- Step 3: build prompt + target → full_ids / label_ids ---- + # Use a dummy image prompt (the actual image tokens will not be used in + # text-only Actor forward, but we need the textual prompt structure) + full_ids_list = [] + tgt_ids_list = [] + + for idx, s in enumerate(real_samples): + # Build the user prompt from history (text-only; images handled at inference) + history = s['history'] + if prior_generator is not None and hasattr(prior_generator, 'get_user_prompt'): + # Use the prior_generator's prompt builder for consistency + valid_actions_hint = [] # not needed for tokenization + user_prompt = prior_generator.get_user_prompt(valid_actions_hint, history) + else: + user_prompt = self.get_user_prompt_image(history=history) + + # Build chat context via tokenizer chat template + prompt_text = self.tokenizer.apply_chat_template( + [ + {"role": "system", "content": self.get_system_prompt_image()}, + {"role": "user", "content": user_prompt}, + ], + tokenize=False, + add_generation_prompt=True, + ) + + target_text = targets_only[idx] + full_text = prompt_text + target_text + + prompt_ids = self.tokenizer.encode(prompt_text, add_special_tokens=False) + full_ids = self.tokenizer.encode(full_text, add_special_tokens=False) + tgt_ids = full_ids[len(prompt_ids):] + + # Truncate prompt if it exceeds prompt_max_len + if len(prompt_ids) > self.prompt_max_len: + prompt_ids = prompt_ids[-self.prompt_max_len:] + full_ids = prompt_ids + tgt_ids + + full_ids_list.append(full_ids) + tgt_ids_list.append(tgt_ids) + + # ---- Step 4: pad and build tensors ---- + inputs = self.tokenizer.pad({"input_ids": full_ids_list}, padding=True, return_tensors="pt") + labels = torch.full_like(inputs.input_ids, -100) + for i, tgt_ids in enumerate(tgt_ids_list): + tgt_len = len(tgt_ids) + labels[i, -tgt_len:] = inputs.input_ids[i, -tgt_len:] + + action_mask_full = (labels != -100).long() + max_tgt_len = max(len(t) for t in tgt_ids_list) + action_mask = action_mask_full[:, -max_tgt_len:] + + # ---- Step 5: compute advantage ---- + target_value_tensor = torch.tensor([s['target_value'] for s in real_samples], dtype=torch.float32) + pred_value_tensor = torch.tensor([s['pred_value'] for s in real_samples], dtype=torch.float32) + advantage = target_value_tensor - pred_value_tensor + + log_status_tmp = {} + + if self.args.advantage_type == "advantage": + log_status_tmp["value_advantage"] = advantage.tolist() + elif self.args.advantage_type == "advantage_batch_norm": + advantage = (advantage - advantage.mean()) / (advantage.std() + 1e-8) + log_status_tmp["value_advantage"] = advantage.tolist() + elif self.args.advantage_type == "advantage_running_norm": + if self.value_normalizer is not None: + advantage_np = advantage.numpy() + advantage_np = self.value_normalizer.normalize_advantages(advantage_np) + advantage = torch.from_numpy(advantage_np) + else: + advantage = (advantage - advantage.mean()) / (advantage.std() + 1e-8) + log_status_tmp["value_advantage"] = advantage.tolist() + else: + log_status_tmp["value_advantage"] = advantage.tolist() + + log_status = [ + {k: log_status_tmp[k][i] for k in log_status_tmp.keys()} + for i in range(len(real_samples)) + ] + + # ---- Step 6: build rollout_logprob ---- + # VL has action-level log-probs (not per-token). We spread the action log-prob + # uniformly across all target tokens so the PPO ratio is correct in expectation. + rollout_logprob = torch.zeros(len(real_samples), max_tgt_len, dtype=torch.float32) + for idx, s in enumerate(real_samples): + tgt_len = len(tgt_ids_list[idx]) + if s['action_logprobs'] is not None and isinstance(s['action_logprobs'], np.ndarray): + # Use MCTS-selected action index to get the correct rollout log-prob. + # Previously used np.max which incorrectly assumed VLM's top choice == MCTS choice. + if s['mcts_action_idx'] is not None and 0 <= s['mcts_action_idx'] < len(s['action_logprobs']): + chosen_logprob = float(s['action_logprobs'][s['mcts_action_idx']]) + else: + # Fallback: use max (legacy behavior, should rarely happen) + chosen_logprob = float(np.max(s['action_logprobs'])) + per_token_lp = chosen_logprob / max(tgt_len, 1) + rollout_logprob[idx, -tgt_len:] = per_token_lp + # else: leave as zero (no rollout log-probs available) + + if self.rank == 0: + _logger.info( + f"[VL Train Samples] Built {len(real_samples)} samples | " + f"advantage mean={advantage.mean().item():.4f} std={advantage.std().item():.4f}" + ) + + return True, (inputs.input_ids, inputs.attention_mask, action_mask, advantage, rollout_logprob, log_status) + + except Exception as e: + import traceback as tb + if self.rank == 0: + _logger.error(f"[VL Train Samples] Error: {e}\n{tb.format_exc()}") + return (False, []) + + def get_llm_output_log(self, wm_train_iter: int, llm_train_iter: int): + """Log LLM/VL output statistics.""" + if self.rank == 0 and len(self.episode_output) > 0: + self._logger.info( + f"[WM Iter {wm_train_iter} | LLM Iter {llm_train_iter}] " + f"Collected {len(self.episode_output)} outputs" + ) + self.episode_output = [] + + +# Backward compatibility: alias to original name +DataProcessor = UnifiedDataProcessor diff --git a/zoo/jericho/priorzero/priorzero_entry.py b/zoo/jericho/priorzero/priorzero_entry.py deleted file mode 100644 index 65337f6f7..000000000 --- a/zoo/jericho/priorzero/priorzero_entry.py +++ /dev/null @@ -1,581 +0,0 @@ -# priorzero_entry.py -""" -[PRIORZERO] Main Training Entry Point - -This module provides the main async training loop for PriorZero. - -Key Features: -- Async training with vLLM integration -- Checkpoint management and recovery -- Comprehensive logging (TensorBoard + file logs) -- Graceful error handling - -Author: PriorZero Team -Date: 2025-01-20 -""" - -import asyncio -import os -import sys -from functools import partial -from pathlib import Path -from typing import Tuple, Optional -# from lzero.entry.utils import log_buffer_memory_usage -# from lzero.policy import visit_count_temperature -# from ding.rl_utils import get_epsilon_greedy_fn - -# ============================================================================== -# [CRITICAL] Ensure local LightZero is used for PriorZero-specific adaptations -# ============================================================================== -from ensure_local_lightzero import ensure_local_lightzero -ensure_local_lightzero() - - -import ray -import torch -import wandb -from ding.config import compile_config -from ding.envs import create_env_manager, get_vec_env_setting -from ding.policy import create_policy -from ding.utils import set_pkg_seed, get_rank -from ding.worker import create_buffer, BaseLearner -from tensorboardX import SummaryWriter -from loguru import logger -from vllm import AsyncLLMEngine -from vllm.engine.arg_utils import AsyncEngineArgs - -# Import PriorZero components -from priorzero_config import get_priorzero_config, get_priorzero_config_for_quick_test -from priorzero_collector import PriorZeroCollector -from priorzero_evaluator import PriorZeroEvaluator -# Import policy to ensure registration happens -import priorzero_policy # noqa: F401 - - -async def train_priorzero( - cfg: dict, - create_cfg: dict, - seed: int = 0, - max_train_iter: int = int(1e6), - max_env_step: Optional[int] = int(1e10), - enable_save: bool = True, -): - """ - [PRIORZERO-MODIFIED] - Main async training function for PriorZero. - - Args: - cfg: Main configuration dictionary - create_cfg: Creation configuration for DI-engine components - seed: Random seed - max_train_iter: Maximum training iterations - enable_save: Whether to save checkpoints - """ - # ================================================================== - # 1. Compile Configuration - # ================================================================== - cfg = compile_config(cfg, seed=seed, auto=True, create_cfg=create_cfg) - - # ================================================================== - # 2. Initialize Ray (for distributed vLLM) - # ================================================================== - # Note: vLLM will initialize Ray internally if needed. - # We skip manual Ray initialization to avoid conflicts with existing clusters. - if ray.is_initialized(): - logger.info(f"✓ Ray already initialized (connected to existing cluster)") - else: - logger.info(f"✓ Ray not initialized - vLLM will handle initialization if needed") - - # ================================================================== - # 3. Create vLLM Engine - # ================================================================== - logger.info("Creating vLLM engine...") - - # [ROBUST FIX] Handle shared GPU environment - # Issue: vLLM V1 engine fails when other processes release GPU memory during init - # Solution: Use alternative initialization method that bypasses V1 checks - import os - - # Note: In vLLM>=0.3.0, worker_use_ray is replaced by distributed_executor_backend - # For single GPU: use "mp" (multiprocessing) - # For multi-GPU: use "ray" if available - tensor_parallel = cfg.policy.llm_policy_cfg.vllm_tensor_parallel_size - distributed_backend = "ray" if tensor_parallel > 1 and ray.is_initialized() else None - - # [ROBUST FIX] Lower GPU memory utilization in shared environment - # This leaves more headroom for memory fluctuations - gpu_mem_util = cfg.policy.llm_policy_cfg.gpu_memory_utilization - if gpu_mem_util > 0.85: - gpu_mem_util = 0.75 # More conservative in shared environment - logger.info(f"✓ Adjusted GPU memory utilization to {gpu_mem_util} for stability") - - # [ROBUST FIX] Use alternative initialization to avoid V1 engine issues - # Set env var BEFORE importing to ensure it takes effect - use_v1_env = os.environ.get('VLLM_USE_V1', None) - if use_v1_env is None: - # Only set if not already set by user - os.environ['VLLM_USE_V1'] = '0' - logger.info("✓ Using vLLM V0 engine for stability in shared GPU environment") - - try: - engine_args = AsyncEngineArgs( - model=cfg.policy.llm_policy_cfg.pretrain_llm_path, - tensor_parallel_size=tensor_parallel, - gpu_memory_utilization=gpu_mem_util, - distributed_executor_backend=distributed_backend, - trust_remote_code=True, - # [ROBUST FIX] Disable prefix caching in shared environment to reduce memory complexity - enable_prefix_caching=False, - # [ROBUST FIX] Disable enforce_eager to avoid memory profiling issues - enforce_eager=False, - ) - vllm_engine = AsyncLLMEngine.from_engine_args(engine_args) - logger.info(f"✓ vLLM Engine created (backend: {distributed_backend or 'default'})") - except (ValueError, RuntimeError) as e: - if "VLLM_USE_V1" in str(e) or "memory profiling" in str(e): - # Fallback: Try without V1 env var - logger.warning(f"⚠️ Initial vLLM initialization failed: {e}") - logger.info("Retrying with alternative configuration...") - if 'VLLM_USE_V1' in os.environ: - del os.environ['VLLM_USE_V1'] - - engine_args = AsyncEngineArgs( - model=cfg.policy.llm_policy_cfg.pretrain_llm_path, - tensor_parallel_size=tensor_parallel, - gpu_memory_utilization=gpu_mem_util * 0.9, # Even more conservative - distributed_executor_backend=distributed_backend, - trust_remote_code=True, - enable_prefix_caching=False, - enforce_eager=True, # Force eager mode as fallback - ) - vllm_engine = AsyncLLMEngine.from_engine_args(engine_args) - logger.info(f"✓ vLLM Engine created with fallback configuration") - else: - raise - - # ================================================================== - # 4. Create Environments - # ================================================================== - logger.info("Creating environments...") - logger.info(f"[DEBUG] Config values: collector_env_num={cfg.env.collector_env_num}, " - f"evaluator_env_num={cfg.env.evaluator_env_num}, " - f"n_evaluator_episode={cfg.env.n_evaluator_episode}") - env_fn, collector_env_cfg, evaluator_env_cfg = get_vec_env_setting(cfg.env) - logger.info(f"[DEBUG] get_vec_env_setting returned: " - f"collector envs={len(collector_env_cfg)}, " - f"evaluator envs={len(evaluator_env_cfg)}") - collector_env = create_env_manager( - cfg.env.manager, - [partial(env_fn, cfg=c) for c in collector_env_cfg] - ) - evaluator_env = create_env_manager( - cfg.env.manager, - [partial(env_fn, cfg=c) for c in evaluator_env_cfg] - ) - - # Seed environments - collector_env.seed(seed) - evaluator_env.seed(seed, dynamic_seed=False) - set_pkg_seed(seed, use_cuda=True) - logger.info(f"✓ Environments created and seeded (seed={seed})") - logger.info(f"[DEBUG] Actual env counts: collector={collector_env.env_num}, " - f"evaluator={evaluator_env.env_num}") - - # ================================================================== - # 5. Create Policy, Buffer, and Components - # ================================================================== - logger.info("Creating policy, buffer, and components...") - - # Create policy (align with UniZero) - policy = create_policy( - cfg.policy, - enable_field=['learn', 'collect', 'eval'] - ) - logger.info("✓ Policy created") - - # Create TensorBoard logger (align with UniZero) - os.makedirs(f'./{cfg.exp_name}/log/', exist_ok=True) - tb_logger = SummaryWriter(os.path.join(f'./{cfg.exp_name}/log/', 'serial')) if get_rank() == 0 else None - logger.info(f"✓ TensorBoard logger: ./{cfg.exp_name}/log/") - - # Create learner (align with UniZero - this sets up policy._logger) - learner = BaseLearner( - cfg.policy.learn.learner, - policy.learn_mode, - tb_logger, - exp_name=cfg.exp_name - ) - logger.info("✓ BaseLearner created") - - # [PRIORZERO-MODIFIED] Create PriorZero-specific replay buffer - # This buffer returns game_segments for LLM training (SFT/RFT) - from lzero.mcts.buffer.game_buffer_priorzero import PriorZeroGameBufferOptimized - replay_buffer = PriorZeroGameBufferOptimized(cfg.policy) - logger.info("✓ PriorZero replay buffer created (with game_segments support)") - - # Create collector - collector = PriorZeroCollector( - env=collector_env, - policy=policy.collect_mode, - tb_logger=tb_logger, - exp_name=cfg.exp_name, - vllm_engine=vllm_engine, - policy_config=cfg.policy, - debug_mode=cfg.get('debug_mode', False), - ) - logger.info("✓ Collector created") - - # Create evaluator - evaluator = PriorZeroEvaluator( - eval_freq=cfg.policy.eval_freq, - n_evaluator_episode=cfg.env.n_evaluator_episode, - stop_value=cfg.env.stop_value, - env=evaluator_env, - policy=policy.eval_mode, - tb_logger=tb_logger, - exp_name=cfg.exp_name, - vllm_engine=vllm_engine, - policy_config=cfg.policy, - ) - logger.info("✓ Evaluator created") - - # Initialize WandB if enabled (PriorZero enhancement) - if cfg.policy.get('use_wandb', True): - if get_rank() == 0: - wandb.init( - project=cfg.policy.get('wandb_project', 'priorzero'), - name=cfg.exp_name, - config=cfg, - tags=['priorzero', 'unizero', 'llm-policy'], - ) - logger.info("✓ WandB initialized") - # Set train iter and env step for policy wandb logging - policy.set_train_iter_env_step(learner.train_iter, collector.envstep) - - # Call learner's before_run hook (align with UniZero) - learner.call_hook('before_run') - - # ================================================================== - # 6. Initialize Async Training Coordinator - # ================================================================== - from async_training_coordinator import AsyncTrainingCoordinator - - coordinator = AsyncTrainingCoordinator( - off_policy_degree=cfg.policy.off_policy_degree, - enable_async_eval=cfg.policy.enable_async_eval, - buffer_size=cfg.policy.replay_buffer_size, - batch_size=cfg.policy.batch_size, - ) - - # ================================================================== - # 7. Main Training Loop - # ================================================================== - logger.info("="*80) - logger.info("Starting PriorZero Training") - logger.info("="*80) - logger.info(f"Experiment: {cfg.exp_name}") - logger.info(f"Max iterations: {max_train_iter}") - logger.info(f"Batch size: {cfg.policy.batch_size}") - logger.info(f"LLM model: {cfg.policy.llm_policy_cfg.pretrain_llm_path}") - logger.info(f"World model layers: {cfg.policy.model.world_model_cfg.num_layers}") - logger.info(f"Off-policy degree: {cfg.policy.off_policy_degree} ({'SYNC' if cfg.policy.off_policy_degree == 0 else 'ASYNC'})") - logger.info(f"Async eval: {cfg.policy.enable_async_eval}") - logger.info("="*80) - - # [ALIGN WITH UNIZERO] Initialize reanalyze-related counters (train_unizero_segment.py line 119-121) - buffer_reanalyze_count = 0 - train_epoch = 0 - reanalyze_batch_size = cfg.policy.reanalyze_batch_size - batch_size = cfg.policy.batch_size - best_eval_reward = -float('inf') - policy_config = cfg.policy - - # Async control variables - collect_task = None - train_task = None - pending_new_data = None # Store collected data waiting to be added to buffer - - try: - while True: - # ================================================================== - # Determine if we're in synchronous or asynchronous mode - # ================================================================== - is_sync_mode = coordinator.is_synchronous - - # ================================================================== - # Evaluation (align with train_unizero_segment.py line 158-162) - # ================================================================== - if learner.train_iter > 0 and evaluator.should_eval(learner.train_iter): - # if learner.train_iter == 0 r evaluator.should_eval(learner.train_iter): - - logger.info(f"\n[Iter {learner.train_iter}] Evaluating...") - - # Define async eval function - async def eval_fn(): - return evaluator.eval( - save_ckpt_fn=learner.save_checkpoint if enable_save else None, - train_iter=learner.train_iter, - envstep=collector.envstep - ) - - # Run eval through coordinator (handles sync/async based on config) - eval_result = await coordinator.run_eval(eval_fn) - - # If sync eval, process result immediately - if not cfg.policy.enable_async_eval and eval_result is not None: - stop, eval_reward_dict = eval_result - mean_reward = eval_reward_dict.get('reward_mean', 0) - logger.info(f" ✓ Evaluation done: reward_mean={mean_reward:.2f}") - - if mean_reward > best_eval_reward: - best_eval_reward = mean_reward - - if stop: - logger.info(f" 🎉 Training converged! (reward >= {cfg.env.stop_value})") - break - else: - logger.info(f" ✓ Async evaluation started in background") - - # ================================================================== - # Collect Data (align with train_unizero_segment.py line 165) - # ================================================================== - collect_kwargs = { - 'temperature': 0.25, - 'epsilon': 0.0 - } - - if is_sync_mode: - # ============================================================ - # SYNCHRONOUS MODE: Original serial execution - # ============================================================ - logger.info(f"\n[Iter {learner.train_iter}] Collecting data...") - - new_data = await collector.collect( - train_iter=learner.train_iter, - policy_kwargs=collect_kwargs - ) - - # Update replay buffer - from lzero.entry.utils import calculate_update_per_collect - update_per_collect = calculate_update_per_collect(cfg, new_data, world_size=1) - - replay_buffer.push_game_segments(new_data) - replay_buffer.remove_oldest_data_to_fit() - buffer_size = replay_buffer.get_num_of_transitions() if hasattr(replay_buffer, 'get_num_of_transitions') else 0 - logger.info(f" ✓ Data collected, buffer size: {buffer_size} transitions") - - else: - # ============================================================ - # ASYNCHRONOUS MODE: Collect can overlap with train - # ============================================================ - # Start or check collect task - if collect_task is None or collect_task.done(): - if coordinator.can_collect(): - logger.info(f"\n[Iter {learner.train_iter}] Starting async collect...") - - # Define async collect function - async def collect_fn(): - return await collector.collect( - train_iter=learner.train_iter, - policy_kwargs=collect_kwargs - ) - - # Start collect task through coordinator - collect_task = asyncio.create_task(coordinator.run_collect(collect_fn)) - else: - logger.debug(f"Collect blocked (lag={coordinator.collect_train_lag}/{coordinator.off_policy_degree})") - - # Check if collect completed - if collect_task is not None and collect_task.done(): - new_data = await collect_task - collect_task = None - - # Store for buffer update - pending_new_data = new_data - logger.info(f" ✓ Async collect completed, data pending buffer update") - - # Update buffer if we have pending data - if pending_new_data is not None: - from lzero.entry.utils import calculate_update_per_collect - update_per_collect = calculate_update_per_collect(cfg, pending_new_data, world_size=1) - - replay_buffer.push_game_segments(pending_new_data) - replay_buffer.remove_oldest_data_to_fit() - buffer_size = replay_buffer.get_num_of_transitions() if hasattr(replay_buffer, 'get_num_of_transitions') else 0 - logger.info(f" ✓ Buffer updated, size: {buffer_size} transitions") - - pending_new_data = None - else: - # No new data yet, use previous update_per_collect or default - update_per_collect = cfg.policy.get('update_per_collect', 10) - - # ============================================================ - # Periodically reanalyze buffer (align with train_unizero_segment.py line 175-186) - # ============================================================ - if cfg.policy.buffer_reanalyze_freq >= 1: - # Reanalyze buffer times in one train_epoch - reanalyze_interval = update_per_collect // cfg.policy.buffer_reanalyze_freq - else: - # Reanalyze buffer each <1/buffer_reanalyze_freq> train_epoch - if train_epoch > 0 and train_epoch % int(1/cfg.policy.buffer_reanalyze_freq) == 0 and replay_buffer.get_num_of_transitions()//cfg.policy.num_unroll_steps > int(reanalyze_batch_size/cfg.policy.reanalyze_partition): - logger.info(f"[Reanalyze] Starting buffer reanalysis...") - replay_buffer.reanalyze_buffer(reanalyze_batch_size, policy) - buffer_reanalyze_count += 1 - logger.info(f" ✓ Buffer reanalyze count: {buffer_reanalyze_count}") - - # ============================================================ - # Training (align with train_unizero_segment.py line 189-221) - # ============================================================ - if collector.envstep > cfg.policy.train_start_after_envsteps: - # Check if there is sufficient data for training - if cfg.policy.sample_type == 'episode': - data_sufficient = replay_buffer.get_num_of_game_segments() > batch_size - else: - data_sufficient = replay_buffer.get_num_of_transitions() > batch_size - - if not data_sufficient: - logger.warning( - f' ⚠ Data in replay_buffer is not sufficient: ' - f'batch_size: {batch_size}, replay_buffer: {replay_buffer}. Continue to collect...' - ) - continue - - logger.info(f"[Iter {learner.train_iter}] Training...") - - # Define training function - async def train_one_batch(): - # Reanalyze buffer during training (align with train_unizero_segment.py line 202-210) - # Note: This is simplified - full reanalyze logic should be per-batch - - # Sample batch - train_data = replay_buffer.sample(batch_size, policy) - train_data.insert(2, learner.train_iter) - - # Train - log_vars = learner.train(train_data, collector.envstep) - - # Update priority if enabled - if cfg.policy.use_priority: - replay_buffer.update_priority(train_data, log_vars[0]['value_priority_orig']) - - return log_vars - - if is_sync_mode: - # Synchronous: train all batches sequentially - for i in range(update_per_collect): - await train_one_batch() - else: - # Asynchronous: train batches while allowing collect to proceed - # We still train sequentially per batch, but collect can run in parallel - if coordinator.can_train(): - # Train one batch through coordinator - await coordinator.run_train(train_one_batch) - else: - logger.debug(f"Train waiting for collect...") - - # Increment epoch counter (align with train_unizero_segment.py line 222) - train_epoch += 1 - - # [FIX] Clear KV cache BEFORE collection to prevent index overflow during MCTS - policy.recompute_pos_emb_diff_and_clear_cache() - - # ============================================================ - # Check stopping criteria (align with train_unizero_segment.py line 226-227) - # ============================================================ - if collector.envstep >= max_env_step or learner.train_iter >= max_train_iter: - logger.info("Stopping condition met, training ends!") - break - - # In async mode, yield to event loop - if not is_sync_mode: - await asyncio.sleep(0.001) - - except KeyboardInterrupt: - logger.warning("\n⚠ Training interrupted by user (Ctrl+C)") - - except Exception as e: - logger.error(f"\n✗ Training error: {e}") - import traceback - traceback.print_exc() - - finally: - # ============================================================ - # Cleanup (align with train_unizero_segment.py line 229) - # ============================================================ - learner.call_hook('after_run') - - # Wait for any pending async eval - if cfg.policy.enable_async_eval: - logger.info("Waiting for async eval to complete...") - await coordinator.wait_for_eval() - - # Print async training statistics - async_stats = coordinator.get_statistics() - logger.info("\n" + "="*80) - logger.info("Async Training Statistics:") - logger.info(f" Mode: {async_stats['mode'].upper()}") - logger.info(f" Collect iterations: {async_stats['collect_count']}") - logger.info(f" Train iterations: {async_stats['train_count']}") - logger.info(f" Final lag: {async_stats['collect_train_lag']}") - if 'collect_avg_time' in async_stats: - logger.info(f" Avg collect time: {async_stats['collect_avg_time']:.2f}s") - if 'train_avg_time' in async_stats: - logger.info(f" Avg train time: {async_stats['train_avg_time']:.2f}s") - if 'eval_avg_time' in async_stats: - logger.info(f" Avg eval time: {async_stats['eval_avg_time']:.2f}s") - logger.info("="*80) - - logger.info("\nCleaning up...") - collector_env.close() - evaluator_env.close() - tb_logger.close() - - logger.info("="*80) - logger.info("Training Complete!") - logger.info(f"Total iterations: {learner.train_iter}") - logger.info(f"Best eval reward: {best_eval_reward:.2f}") - logger.info("="*80) - - return policy - - -def main(): - """ - Main entry point with argument parsing. - """ - import argparse - - parser = argparse.ArgumentParser(description='PriorZero Training') - parser.add_argument('--env_id', type=str, default='zork1.z5', help='Jericho game ID') - parser.add_argument('--seed', type=int, default=0, help='Random seed') - parser.add_argument('--max_iter', type=int, default=int(1e6), help='Max training iterations') - parser.add_argument('--quick_test', action='store_true', help='Use quick test config') - parser.add_argument('--no_save', action='store_true', help='Disable checkpoint saving') - parser.add_argument('--debug', action='store_true', help='Enable detailed debug logging (obs, action, LLM output)') - - args = parser.parse_args() - - # args.quick_test = True # ONLY FOR DEBUG - - # Get configuration - if args.quick_test: - logger.info("Using quick test configuration") - main_cfg, create_cfg = get_priorzero_config_for_quick_test(args.env_id, args.seed, debug_mode=args.debug) - else: - main_cfg, create_cfg = get_priorzero_config(args.env_id, args.seed, debug_mode=args.debug) - - # Run training - asyncio.run(train_priorzero( - main_cfg, - create_cfg, - seed=args.seed, - max_train_iter=args.max_iter, - enable_save=not args.no_save - )) - - -if __name__ == "__main__": - import os - # Disable tokenizer parallelism to prevent multi-process conflicts - os.environ['TOKENIZERS_PARALLELISM'] = 'false' - main() diff --git a/zoo/jericho/priorzero/priorzero_entry_unified.py b/zoo/jericho/priorzero/priorzero_entry_unified.py new file mode 100644 index 000000000..9745bec0d --- /dev/null +++ b/zoo/jericho/priorzero/priorzero_entry_unified.py @@ -0,0 +1,754 @@ +""" +Complete PriorZero Entry with VL (Vision-Language) Support + +This is the COMPLETE implementation with full training loop. +Supports both text (LLM) and image (VL) inputs. +""" +import sys +import os +from pathlib import Path + +# Add project root to path +current_file_path = Path(__file__).resolve() +project_root = current_file_path.parents[3] +if str(project_root) not in sys.path: + print(f"[SYSTEM] Inserting project root to sys.path: {project_root}") + sys.path.insert(0, str(project_root)) + +# Add src/ directory to path so that modules like strategy, models, vllm_utils, utils can be found +src_dir = current_file_path.parent / 'src' +if str(src_dir) not in sys.path: + print(f"[SYSTEM] Inserting src dir to sys.path: {src_dir}") + sys.path.insert(0, str(src_dir)) + +# Add priorzero/ directory itself so sibling modules (prior_generator, vl_config, etc.) can be found +priorzero_dir = str(current_file_path.parent) +if priorzero_dir not in sys.path: + print(f"[SYSTEM] Inserting priorzero dir to sys.path: {priorzero_dir}") + sys.path.insert(0, priorzero_dir) + +import argparse +from functools import partial +from typing import Tuple, Optional, List + +import torch +import torch.distributed as dist +from ding.config import compile_config +from ding.envs import create_env_manager, get_vec_env_setting +from ding.policy import create_policy +from ding.utils import set_pkg_seed, get_rank, get_world_size +from ding.worker import BaseLearner +from tensorboardX import SummaryWriter +from loguru import logger + +from lzero.mcts.buffer.game_buffer_priorzero import PriorZeroGameBufferOptimized +from lzero.entry.utils import calculate_update_per_collect + + +def all_gather_cmd(world_size, obj) -> List: + """Gather command from all ranks.""" + if world_size <= 1: + return [obj] + lst = [None] * dist.get_world_size() + dist.all_gather_object(lst, obj) + return lst + + +def prepare_common_components(rank, cfg, create_cfg, seed): + """Prepare components common to both LLM and VL.""" + cfg = compile_config(cfg, seed=seed, auto=True, create_cfg=create_cfg) + + # Create environments + env_fn, collector_env_cfg, evaluator_env_cfg = get_vec_env_setting(cfg.env) + collector_env = create_env_manager(cfg.env.manager, [partial(env_fn, cfg=c) for c in collector_env_cfg]) + evaluator_env = create_env_manager(cfg.env.manager, [partial(env_fn, cfg=c) for c in evaluator_env_cfg]) + + collector_env.seed(seed) + evaluator_env.seed(seed, dynamic_seed=False) + + # Create policy + policy = create_policy(cfg.policy, enable_field=['learn', 'collect', 'eval'], exp_name=cfg.exp_name) + logger.info(f"[Rank {rank}] Policy created") + + # Create logger and learner + os.makedirs(f'./{cfg.exp_name}/log/', exist_ok=True) + tb_logger = SummaryWriter(os.path.join(f'./{cfg.exp_name}/log/', 'serial')) if rank == 0 else None + logger.info(f"[Rank {rank}] TensorBoard logger: ./{cfg.exp_name}/log/") + + learner = BaseLearner(cfg.policy.learn.learner, policy.learn_mode, tb_logger, exp_name=cfg.exp_name) + logger.info(f"[Rank {rank}] BaseLearner created") + + # Create replay buffer + replay_buffer = PriorZeroGameBufferOptimized(cfg.policy) + logger.info(f"[Rank {rank}] PriorZero replay buffer created") + + return cfg, collector_env, evaluator_env, policy, learner, replay_buffer, tb_logger + + +def prepare_llm_components(rank, cfg, llm_cfg, strategy, collector_env, evaluator_env, policy, tb_logger, seed): + """Prepare LLM-specific components for text input.""" + from utils import Profiler, dump_dataclass_cfg_py + from models.actor import PolicyModel, ReferenceModel + from vllm_utils.vllm_engine import create_vllm_engine + from priorzero_datafactory_unified import UnifiedDataProcessor + from priorzero_trainer import PriorZeroLLMTrainer + from priorzero_collector_unified import PriorZeroCollector + from priorzero_evaluator import PriorZeroEvaluator + from prior_generator import LLMPriorGenerator + + prof = Profiler(log_interval=10, stats_file=f'./{cfg.exp_name}/log/profiler.txt', enable_profile=False) + + if rank == 0: + dump_dataclass_cfg_py(llm_cfg, path=f"{cfg.exp_name}/llm_cfg.py") + llm_cfg.save_path = f'./{cfg.exp_name}/llm_ckpt/' + + logger.info(f"[Rank {rank}] Initializing LLM components...") + set_pkg_seed(seed + rank, use_cuda=True) + + # Reference model + ref_model = ReferenceModel(strategy=strategy, pretrain=llm_cfg.model_name_or_path) if llm_cfg.rft_kl_coef > 0 else None + + # vLLM engine + vllm_engine = create_vllm_engine( + tensor_parallel_size=llm_cfg.vllm_tensor_parallel_size, + pretrain=llm_cfg.model_name_or_path, + enable_prefix_caching=llm_cfg.enable_prefix_caching, + max_model_len=llm_cfg.prompt_max_len + llm_cfg.generate_max_len, + gpu_memory_utilization=llm_cfg.gpu_memory_utilization, + vllm_enable_sleep=llm_cfg.vllm_enable_sleep, + ) + logger.info(f'[Rank {rank}] vLLM engine created') + + # Data processor + world_size = getattr(strategy, "world_size", 1) + data_processor = UnifiedDataProcessor( + rank=rank, + world_size=world_size, + vllm_engine=vllm_engine, + strategy=strategy, + model_path=llm_cfg.model_name_or_path, + exp_name=cfg.exp_name if rank == 0 else None, + obs_type='text', + ) + + # Policy model + policy_model = PolicyModel( + strategy=strategy, + pretrain=llm_cfg.model_name_or_path, + vllm_engine=vllm_engine, + max_steps=llm_cfg.max_steps + ) + + # Trainer + trainer = PriorZeroLLMTrainer( + cfg=llm_cfg, + pretrain=llm_cfg.model_name_or_path, + strategy=strategy, + vllm_engine=vllm_engine, + policy_model=policy_model, + reference_model=ref_model, + exp_name=cfg.exp_name if rank == 0 else None, + tb_logger=tb_logger if rank == 0 else None, + llm_save_freq=llm_cfg.llm_save_freq + ) + + # Prior generator + prior_generator = LLMPriorGenerator( + vllm_engine=vllm_engine, + data_processor=data_processor, + model_name=llm_cfg.model_name_or_path, + use_cot=llm_cfg.use_cot, + ) + + # Collector + collector = PriorZeroCollector( + env=collector_env, + policy=policy.collect_mode, + llm_config=llm_cfg, + tb_logger=tb_logger, + exp_name=cfg.exp_name, + policy_config=cfg.policy, + data_processor=data_processor, + prior_generator=prior_generator, + obs_type='text', + ) + collector.prof = prof + + # Evaluator + evaluator = PriorZeroEvaluator( + eval_freq=cfg.policy.eval_freq, + n_evaluator_episode=cfg.env.n_evaluator_episode, + stop_value=cfg.env.stop_value, + env=evaluator_env, + policy=policy.eval_mode, + tb_logger=tb_logger, + exp_name=cfg.exp_name, + policy_config=cfg.policy, + llm_config=llm_cfg, + data_processor=data_processor, + ) + + logger.info(f"[Rank {rank}] ✓ LLM components initialized") + + return { + 'prior_generator': prior_generator, + 'vllm_engine': vllm_engine, + 'policy_model': policy_model, + 'ref_model': ref_model, + 'trainer': trainer, + 'data_processor': data_processor, + 'collector': collector, + 'evaluator': evaluator, + 'prof': prof, + } + + +def prepare_vl_components(rank, cfg, vl_cfg, strategy, collector_env, evaluator_env, policy, tb_logger, seed): + """Prepare VL-specific components for image input.""" + from utils import Profiler, dump_dataclass_cfg_py + from models.actor import PolicyModel, ReferenceModel + from vl_engine import create_vl_engine + from priorzero_datafactory_unified import UnifiedDataProcessor + from priorzero_trainer import PriorZeroLLMTrainer # Can reuse for VL + from priorzero_collector_unified import PriorZeroCollector + from priorzero_evaluator import PriorZeroEvaluator + from prior_generator import VLPriorGenerator + + prof = Profiler(log_interval=10, stats_file=f'./{cfg.exp_name}/log/profiler.txt', enable_profile=False) + + if rank == 0: + dump_dataclass_cfg_py(vl_cfg, path=f"{cfg.exp_name}/vl_cfg.py") + vl_cfg.save_path = f'./{cfg.exp_name}/vl_ckpt/' + + logger.info(f"[Rank {rank}] Initializing VL components...") + set_pkg_seed(seed + rank, use_cuda=True) + + # Reference model + ref_model = ReferenceModel(strategy=strategy, pretrain=vl_cfg.model_name_or_path) if vl_cfg.rft_kl_coef > 0 else None + + # VL engine + # Determine limit_mm_per_prompt based on vlm_image_mode + vlm_image_mode = getattr(vl_cfg, 'vlm_image_mode', 'current_only') + if vlm_image_mode == "current_only": + limit_mm_per_prompt = {"image": 1} + else: + # first_and_current or all_history: need up to history_length + 1 images + limit_mm_per_prompt = {"image": vl_cfg.history_length + 1} + + vl_engine = create_vl_engine( + model_name=vl_cfg.vl_model_type, + model_path=vl_cfg.model_name_or_path, + tensor_parallel_size=vl_cfg.tensor_parallel_size, + gpu_memory_utilization=vl_cfg.gpu_memory_utilization, + max_model_len=vl_cfg.prompt_max_len + vl_cfg.generate_max_len, + limit_mm_per_prompt=limit_mm_per_prompt, + ) + logger.info(f'[Rank {rank}] VL engine created: {vl_cfg.vl_model_type}') + + # Data processor + world_size = getattr(strategy, "world_size", 1) + data_processor = UnifiedDataProcessor( + rank=rank, + world_size=world_size, + vllm_engine=vl_engine, + strategy=strategy, + model_path=vl_cfg.model_name_or_path, + exp_name=cfg.exp_name if rank == 0 else None, + obs_type='image', + ) + + # Policy model + policy_model = PolicyModel( + strategy=strategy, + pretrain=vl_cfg.model_name_or_path, + vllm_engine=vl_engine, + max_steps=vl_cfg.max_steps + ) + + # Trainer + trainer = PriorZeroLLMTrainer( + cfg=vl_cfg, + pretrain=vl_cfg.model_name_or_path, + strategy=strategy, + vllm_engine=vl_engine, + policy_model=policy_model, + reference_model=ref_model, + exp_name=cfg.exp_name if rank == 0 else None, + tb_logger=tb_logger if rank == 0 else None, + instance_name="vl_ppo", + llm_save_freq=vl_cfg.vl_save_freq + ) + + # Prior generator + prior_generator = VLPriorGenerator( + vl_engine=vl_engine, + model_name=vl_cfg.model_name_or_path, + use_cot=vl_cfg.use_cot, + game_description=getattr(vl_cfg, 'game_description', ''), + vlm_image_mode=vlm_image_mode, + prompt_style=getattr(vl_cfg, 'prompt_style', 'concise'), + logprob_extraction_mode=getattr(vl_cfg, 'logprob_extraction_mode', 'approximate'), + ) + + # Collector + collector = PriorZeroCollector( + env=collector_env, + policy=policy.collect_mode, + llm_config=vl_cfg, + tb_logger=tb_logger, + exp_name=cfg.exp_name, + policy_config=cfg.policy, + data_processor=data_processor, + prior_generator=prior_generator, + obs_type='image', + env_id=cfg.env.env_id, # Pass env_id for action mapping + ) + collector.prof = prof + + # Evaluator + evaluator = PriorZeroEvaluator( + eval_freq=cfg.policy.eval_freq, + n_evaluator_episode=cfg.env.n_evaluator_episode, + stop_value=cfg.env.stop_value, + env=evaluator_env, + policy=policy.eval_mode, + tb_logger=tb_logger, + exp_name=cfg.exp_name, + policy_config=cfg.policy, + llm_config=vl_cfg, + data_processor=data_processor, + prior_generator=prior_generator, + obs_type='image', + env_id=cfg.env.env_id, + ) + + logger.info(f"[Rank {rank}] ✓ VL components initialized") + + return { + 'prior_generator': prior_generator, + 'vl_engine': vl_engine, + 'policy_model': policy_model, + 'ref_model': ref_model, + 'trainer': trainer, + 'data_processor': data_processor, + 'collector': collector, + 'evaluator': evaluator, + 'prof': prof, + } + + +def train_unified( + cfg: dict, + create_cfg: dict, + prior_cfg, # LLM or VL config + seed: int = 0, + max_train_iter: int = int(1e6), + max_env_step: Optional[int] = int(1e10), + enable_profile: bool = False, + is_text_input: bool = True, +): + """ + Unified training function supporting both LLM and VL. + + Args: + cfg: Main configuration + create_cfg: Creation configuration + prior_cfg: LLM or VL configuration + seed: Random seed + max_train_iter: Maximum training iterations + max_env_step: Maximum environment steps + enable_profile: Whether to enable profiling + is_text_input: Whether using text input (True) or image input (False) + """ + rank = int(os.environ.get("RANK", "0")) + + # Initialize strategy + from strategy.deepspeed import get_strategy, torch_dist_barrier_and_cuda_sync + strategy = get_strategy(prior_cfg) + strategy.print(prior_cfg) + strategy.setup_distributed() + world_size = getattr(strategy, "world_size", 1) + + # Prepare common components + cfg, collector_env, evaluator_env, policy, learner, replay_buffer, tb_logger = prepare_common_components( + rank, cfg, create_cfg, seed + ) + batch_size = cfg.policy.batch_size + + # Prepare input-specific components + if is_text_input: + components = prepare_llm_components( + rank, cfg, prior_cfg, strategy, collector_env, evaluator_env, policy, tb_logger, seed + ) + engine_name = "vLLM" + else: + components = prepare_vl_components( + rank, cfg, prior_cfg, strategy, collector_env, evaluator_env, policy, tb_logger, seed + ) + engine_name = "VL" + + # Extract components + prior_engine = components['vllm_engine'] if is_text_input else components['vl_engine'] + policy_model = components['policy_model'] + trainer = components['trainer'] + data_processor = components['data_processor'] + collector = components['collector'] + evaluator = components['evaluator'] + prof = components['prof'] + + # Set llm_cfg on policy so _forward_eval/_forward_collect can access it + policy.llm_cfg = prior_cfg + + torch_dist_barrier_and_cuda_sync() + learner.call_hook('before_run') + + logger.info(f"[Rank {rank}] Starting training loop with {engine_name} prior...") + + # Validate VL config consistency (e.g. enable_rft + vl_fixed conflict) + if not is_text_input and hasattr(prior_cfg, 'validate'): + prior_cfg.validate() + + # ========================================================================= + # Alternating Training Schedule Setup (aligned with sync_ddp) + # ========================================================================= + train_schedule = prior_cfg.train_schedule + train_alternate = train_schedule["alternate"] + enable_world_model = prior_cfg.enable_world_model + enable_rft = prior_cfg.enable_rft + + if train_alternate: + current_phase = train_schedule["start_phase"] + last_wm_train_iter = 0 + last_llm_train_iter = 0 + else: + current_phase = None + + # ========================================================================= + # Main Training Loop + # ========================================================================= + while True: + if collector.envstep >= max_env_step or learner.train_iter >= max_train_iter: + break + + # Periodic loop status log (every 500 envsteps, rank 0 only) + if rank == 0 and collector.envstep % 500 < 10: + phase_str = current_phase if train_alternate else "joint" + logger.info( + f"[Loop Status] envstep={collector.envstep}, wm_iter={learner.train_iter}, " + f"llm_iter={trainer.global_step if hasattr(trainer, 'global_step') else 'N/A'}, " + f"phase={phase_str}, enable_rft={enable_rft}, enable_wm={enable_world_model}" + ) + + cmd = 0 + priorzero_batch = None + + # Evaluation + # if learner.train_iter == 0 or evaluator.should_eval(learner.train_iter): + if learner.train_iter > 0 and evaluator.should_eval(learner.train_iter): + logger.info(f"\n[Rank {rank}: Iter {learner.train_iter}] Evaluating...") + if prior_cfg.vllm_enable_sleep and prior_engine is not None: + prior_engine.wake_up() + stop, reward = evaluator.eval( + train_iter=learner.train_iter, + envstep=collector.envstep + ) + if prior_cfg.vllm_enable_sleep and prior_engine is not None: + prior_engine.sleep() + torch.cuda.empty_cache() + + # Wake up engine + if prior_cfg.vllm_enable_sleep and prior_engine is not None: + prior_engine.wake_up() + + # Data collection + with prof.block("collect", rank=rank): + new_data = collector.collect( + train_iter=learner.train_iter, + policy_kwargs={'temperature': 0.25, 'epsilon': 0.0} + ) + + # Log output based on input type + if is_text_input: + data_processor.get_llm_output_log( + wm_train_iter=learner.train_iter, + llm_train_iter=policy_model.train_iter + ) + else: + # VL: use prior_generator's log method + prior_generator = components.get('prior_generator') + if prior_generator and hasattr(prior_generator, 'get_vl_output_log'): + prior_generator.get_vl_output_log( + wm_train_iter=learner.train_iter, + vl_train_iter=policy_model.train_iter + ) + + # Sleep engine + if prior_cfg.vllm_enable_sleep and prior_engine is not None: + prior_engine.sleep() + + # Push to replay buffer + replay_buffer.push_game_segments(new_data) + replay_buffer.remove_oldest_data_to_fit() + + num_of_transitions = replay_buffer.get_num_of_transitions() + new_num_of_transitions = num_of_transitions - replay_buffer.last_pos_in_transition + + logger.info( + f"[Data Collection] Rank {rank} | " + f"Total transitions: {num_of_transitions} | " + f"New transitions: {new_num_of_transitions}" + ) + + # TB logging for collect metrics + if tb_logger is not None: + tb_logger.add_scalar('collect/num_transitions', num_of_transitions, collector.envstep) + tb_logger.add_scalar('collect/new_transitions', new_num_of_transitions, collector.envstep) + + torch_dist_barrier_and_cuda_sync() + + # ===================================================================== + # World Model Training (gated by schedule) + # ===================================================================== + if enable_world_model and (not train_alternate or current_phase == "wm"): + if not (num_of_transitions > batch_size): + logger.warning( + f'[WM Training] Data insufficient: batch_size={batch_size}, buffer={num_of_transitions}' + ) + cmd = 0 + else: + cmd = 1 + + if min(all_gather_cmd(world_size=world_size, obj=cmd)) == 0: + continue + + update_per_collect = calculate_update_per_collect(cfg, new_data, world_size=world_size) + logger.info( + f"[WM Training] Rank {rank} | Iter {learner.train_iter} | " + f"Updates: {update_per_collect}" + ) + + for i in range(update_per_collect): + with prof.block("train_world_model", rank=rank): + train_data = replay_buffer.sample(batch_size, policy) + train_data.append(learner.train_iter) + log_vars = learner.train(train_data, collector.envstep) + if cfg.policy.use_priority: + replay_buffer.update_priority(train_data, log_vars[0]['value_priority_orig']) + + policy.recompute_pos_emb_diff_and_clear_cache() + + # TB logging for WM training + # NOTE: DI-engine's BaseLearner already logs all _monitor_vars_learn() metrics + # under "learner_iter/" prefix (averaged over log_show_after_iter). + # We only log the phase-tracking scalar here; per-metric logging is handled by the learner. + if tb_logger is not None: + tb_logger.add_scalar('train/wm_train_iter', learner.train_iter, collector.envstep) + + # Phase switching: WM -> LLM/VL + if train_alternate and learner.train_iter - last_wm_train_iter >= train_schedule["wm_update_iters"]: + current_phase = "llm" + last_wm_train_iter = learner.train_iter + replay_buffer.mark_latest_transitions_consumed() + logger.info(f"[WM Training][Rank {rank}] Switching to {'VL' if not is_text_input else 'LLM'} training phase at wm iter: {learner.train_iter}") + continue + + # ===================================================================== + # LLM/VL Training (gated by schedule) + # ===================================================================== + if enable_rft and (not train_alternate or current_phase == "llm"): + new_num_of_transitions = replay_buffer.get_num_of_transitions() - replay_buffer.last_pos_in_transition + logger.info( + f"[{engine_name} Training] Rank {rank} | " + f"Total transitions: {num_of_transitions} | " + f"New transitions: {new_num_of_transitions}" + ) + + with prof.block("fetch_latest_batch", rank=rank): + priorzero_batch = replay_buffer.fetch_latest_batch(batch_size=-1, policy=policy) + torch.cuda.empty_cache() + + with prof.block("train_prior_model", rank=rank): + llm_need_sample_cnt = prior_cfg.train_batch_size * prior_cfg.max_rollout_staleness // world_size + flag, train_samples = data_processor.make_llm_train_samples( + priorzero_batch, ddp=True, max_samples=llm_need_sample_cnt, + prior_generator=components.get('prior_generator') if not is_text_input else None, + ) + + if not flag: + local_llm_ready = 0 + else: + local_llm_ready = 1 + + gathered_llm_ready = all_gather_cmd(world_size=world_size, obj=local_llm_ready) + + if min(gathered_llm_ready) == 0: + logger.info( + f"[Rank {rank}] Skip {engine_name} training: not all ranks ready. " + f"ready_flags={gathered_llm_ready}, local={local_llm_ready}, required={llm_need_sample_cnt}, got={len(train_samples)}" + ) + continue + + trainer.train_batch(train_samples, collect_env_steps=collector.envstep) + replay_buffer.mark_latest_transitions_consumed() + + torch_dist_barrier_and_cuda_sync() + + # Phase switching: LLM/VL -> WM + if train_alternate and trainer.global_step - last_llm_train_iter >= train_schedule["llm_update_iters"]: + current_phase = "wm" + last_llm_train_iter = trainer.global_step + if data_processor.value_normalizer is not None: + data_processor.value_normalizer.clear() + logger.info(f"[Rank {rank}] Switching to World Model training phase at {engine_name} iter: {trainer.global_step}") + + # Safety fallback: if RFT is disabled but the alternating scheduler + # switched to "llm" phase, immediately fall back to WM so the loop + # does not spin forever doing only data collection. + if not enable_rft and train_alternate and current_phase == "llm": + current_phase = "wm" + logger.info( + f"[Rank {rank}] enable_rft=False, auto-switching from '{engine_name}' phase back to WM phase" + ) + + logger.info(f"[Rank {rank}] Training completed!") + + +def main(): + """Main entry point.""" + parser = argparse.ArgumentParser(description='PriorZero with VL Support') + + # Common arguments + parser.add_argument('--input_type', type=str, required=True, choices=['text', 'image']) + parser.add_argument('--env_id', type=str, required=True) + parser.add_argument('--seed', type=int, default=0) + parser.add_argument('--max_iter', type=int, default=int(1e6)) + parser.add_argument('--quick_test', action='store_true', default=False) + parser.add_argument('--enable_profile', action='store_true', default=False) + + # Text-specific + parser.add_argument('--llm_model', type=str, default='qwen2.5-1.5b') + + # Shared LLM/VL arguments + parser.add_argument('--use_cot', action='store_true', default=True, + help='Enable Chain-of-Thought reasoning (default: True)') + parser.add_argument('--no_cot', action='store_true', default=False, + help='Disable Chain-of-Thought reasoning') + parser.add_argument('--cot_weight', type=float, default=0.1, + help='Weight for CoT prefix tokens in loss (default: 0.1)') + parser.add_argument('--vl_fixed', action='store_true', default=True, + help='Freeze VL model (inference only, no VL training) (default: True)') + parser.add_argument('--no_vl_fixed', action='store_true', default=False, + help='Enable VL training (unfreeze)') + parser.add_argument('--mcts_mode', type=str, default='llm_plus_wm_logits', + choices=['llm_logits', 'wm_logits', 'llm_plus_wm_logits'], + help='MCTS root logits mode (default: llm_plus_wm_logits)') + + # Image-specific + parser.add_argument('--vl_model', type=str, default='Qwen2.5-VL-7b') + parser.add_argument('--use_prior', action='store_true', default=True) + parser.add_argument('--vlm_image_mode', type=str, default='current_only', + choices=['current_only', 'first_and_current', 'all_history'], + help='VLM image mode: how many images to send to VL model (default: current_only)') + parser.add_argument('--prompt_style', type=str, default='legacy', + choices=['concise', 'legacy'], + help='Prompt style: concise (shorter, better for small VLMs) or legacy (verbose)') + + args = parser.parse_args() + + # Resolve --no_xxx flags (explicit --no_cot / --no_vl_fixed override defaults) + args.use_cot = not args.no_cot + args.vl_fixed = not args.no_vl_fixed + + print(f"\n{'='*80}") + print(f"PriorZero Training with {'LLM' if args.input_type == 'text' else 'VL'} Prior") + print(f"{'='*80}") + print(f"Input Type: {args.input_type}") + print(f"Environment: {args.env_id}") + print(f"Seed: {args.seed}") + print(f"Quick Test: {args.quick_test}") + print(f"Use CoT: {args.use_cot}") + if args.input_type == 'image': + print(f"VL Fixed: {args.vl_fixed}") + print(f"MCTS Mode: {args.mcts_mode}") + print(f"VLM Image Mode: {args.vlm_image_mode}") + print(f"{'='*80}\n") + + if args.input_type == 'text': + from priorzero_config import get_priorzero_config, get_priorzero_debug_config + + if args.quick_test: + main_cfg, create_cfg, llm_cfg = get_priorzero_debug_config( + args.env_id, args.seed, use_cot=args.use_cot, + exp_name=f'data_priorzero_complete/text_{args.env_id}_seed{args.seed}', + model_key=args.llm_model, + ) + else: + main_cfg, create_cfg, llm_cfg = get_priorzero_config( + args.env_id, args.seed, use_cot=args.use_cot, + exp_name=f'data_priorzero_complete/text_{args.env_id}_seed{args.seed}', + model_key=args.llm_model, + multi_gpu=True + ) + + train_unified( + main_cfg, create_cfg, llm_cfg, + seed=args.seed, + max_train_iter=args.max_iter, + enable_profile=args.enable_profile, + is_text_input=True, + ) + + else: + from vl_config import get_priorzero_vl_config + + # Build a clean env short name: strip common suffixes + env_short = args.env_id + for suffix in ['NoFrameskip-v4', '-v2', '-v1', '-v0', '-v5']: + if env_short.endswith(suffix): + env_short = env_short[:-len(suffix)] + break + + from datetime import datetime + timestamp = datetime.now().strftime('%y%m%d_%H%M%S') + cot_tag = f"cot{args.cot_weight}" if args.use_cot else "noCot" + fixed_tag = "vlFixed" if args.vl_fixed else "vlTrain" + exp_name = ( + f'data_priorzero_complete/' + f'{env_short}_{args.vl_model}_{fixed_tag}/' + f'{cot_tag}_mcts_{args.mcts_mode}_img_{args.vlm_image_mode}/' + f'seed{args.seed}_{timestamp}' + ) + + main_cfg, create_cfg, vl_cfg = get_priorzero_vl_config( + args.env_id, args.seed, + exp_name=exp_name, + vl_model_key=args.vl_model, + use_prior=args.use_prior, + multi_gpu=int(os.environ.get('WORLD_SIZE', '1')) > 1, + quick_test=args.quick_test, + ) + + # Apply CLI overrides to vl_cfg + if vl_cfg is not None: + vl_cfg.use_cot = args.use_cot + vl_cfg.cot_weight = args.cot_weight + vl_cfg.vl_fixed = args.vl_fixed + vl_cfg.mcts_root_logits_dict.mode = args.mcts_mode + vl_cfg.vlm_image_mode = args.vlm_image_mode + vl_cfg.prompt_style = args.prompt_style + # Ensure consistency: vl_fixed=True → disable PPO training + if vl_cfg.vl_fixed: + vl_cfg.enable_rft = False + + train_unified( + main_cfg, create_cfg, vl_cfg, + seed=args.seed, + max_train_iter=args.max_iter, + enable_profile=args.enable_profile, + is_text_input=False, + ) + + +if __name__ == "__main__": + os.environ['TOKENIZERS_PARALLELISM'] = 'false' + main() diff --git a/zoo/jericho/priorzero/priorzero_evaluator.py b/zoo/jericho/priorzero/priorzero_evaluator.py deleted file mode 100644 index c8a25d0f9..000000000 --- a/zoo/jericho/priorzero/priorzero_evaluator.py +++ /dev/null @@ -1,55 +0,0 @@ -# priorzero_evaluator.py -""" -[PRIORZERO] PriorZero Evaluator - -Simple evaluator that inherits from MuZeroEvaluator. -Since the policy already integrates LLM priors in its _forward_collect method, -the evaluator can use the parent implementation directly. - -Author: PriorZero Team -Date: 2025-01-20 -""" - -from typing import Optional - -from ding.worker.collector.base_serial_evaluator import SERIAL_EVALUATOR_REGISTRY -from lzero.worker.muzero_evaluator import MuZeroEvaluator as OriginalEvaluator -from vllm import AsyncLLMEngine - - -@SERIAL_EVALUATOR_REGISTRY.register('priorzero', force_overwrite=True) -class PriorZeroEvaluator(OriginalEvaluator): - """ - [PRIORZERO-MODIFIED] - Evaluator for PriorZero. - - Since the PriorZero policy already integrates LLM priors in its - _forward_collect method, this evaluator simply inherits all - functionality from MuZeroEvaluator. - - The vLLM engine is passed for potential future enhancements - (e.g., comparative evaluation with/without LLM priors). - """ - - def __init__( - self, - vllm_engine: Optional[AsyncLLMEngine] = None, - **kwargs - ): - """ - Initialize PriorZeroEvaluator. - - Args: - vllm_engine: vLLM async engine (optional, for future use) - **kwargs: Arguments for parent MuZeroEvaluator - """ - super().__init__(**kwargs) - self.vllm_engine = vllm_engine - - if vllm_engine is not None: - self._logger.info("✓ PriorZeroEvaluator initialized with vLLM engine") - else: - self._logger.info("✓ PriorZeroEvaluator initialized (no vLLM engine)") - - # All other methods are inherited from MuZeroEvaluator - # The policy's _forward_collect already handles LLM prior integration diff --git a/zoo/jericho/priorzero/priorzero_orz_complete.py b/zoo/jericho/priorzero/priorzero_orz_complete.py deleted file mode 100644 index f0daf5958..000000000 --- a/zoo/jericho/priorzero/priorzero_orz_complete.py +++ /dev/null @@ -1,965 +0,0 @@ -""" -PriorZero-ORZ Complete Integration -完整可执行版本 with ORZ RayPPOTrainer - -This version includes: -1. Fixed vLLM None handling -2. Fixed asyncio scope issue -3. Complete ORZ RayPPOTrainer integration -4. Robust error handling - -Usage: - DEBUG_MODE=True python -m zoo.jericho.priorzero.priorzero_orz_complete - -Author: PriorZero Team -Date: 2025-10-21 -""" - -import asyncio -import os -import sys -import re -from pathlib import Path -from functools import partial -from typing import Optional, List, Dict, Any, Callable, Awaitable, Tuple -import time -import json - -# ============================================================================== -# Ensure local LightZero is used -# ============================================================================== -from ensure_local_lightzero import ensure_local_lightzero -ensure_local_lightzero() - -import torch -import numpy as np -from ding.config import compile_config -from ding.envs import create_env_manager, get_vec_env_setting -from ding.policy import create_policy -from ding.utils import set_pkg_seed, get_rank -from ding.worker import BaseLearner -from tensorboardX import SummaryWriter -from loguru import logger - -# PriorZero imports -from priorzero_config import get_priorzero_config_for_quick_test, get_priorzero_config -from priorzero_collector import PriorZeroCollector -from priorzero_evaluator import PriorZeroEvaluator -import priorzero_policy # noqa: F401 -from lzero.mcts.buffer.game_buffer_priorzero import PriorZeroGameBufferOptimized - -# vLLM imports (optional) -try: - from vllm import AsyncLLMEngine - from vllm.engine.arg_utils import AsyncEngineArgs - VLLM_AVAILABLE = True -except ImportError: - VLLM_AVAILABLE = False - logger.warning("vLLM not available - LLM inference will be disabled") - -# Try to import ORZ -ORZ_AVAILABLE = False -ORZ_PATH = Path("/mnt/nfs/zhangjinouwen/puyuan/Open-Reasoner-Zero") - -try: - if ORZ_PATH.exists() and str(ORZ_PATH) not in sys.path: - sys.path.insert(0, str(ORZ_PATH)) - - from orz.ppo import RayPPOTrainer, PromptDataset - from orz.exps.examples.ppo.ppo_base_exp import BasePPOExp, BasePPOExpConfig - from orz.ppo.utils import get_strategy - from transformers import AutoTokenizer - import ray - ORZ_AVAILABLE = True - logger.info("✅ ORZ available - will use ORZ RayPPOTrainer for LLM training") -except ImportError as e: - logger.warning(f"⚠️ ORZ not available ({e}) - will use PriorZero's built-in LLM training") - - -# ============================================================================== -# Configuration -# ============================================================================== - -DEBUG_MODE = os.environ.get("DEBUG_MODE", "False") == "True" - - -class HybridTrainingConfig: - """ - Hybrid training configuration combining PriorZero and ORZ settings. - """ - def __init__(self): - # Get base PriorZero config - if DEBUG_MODE: - self.priorzero_cfg, self.priorzero_create_cfg = get_priorzero_config_for_quick_test( - env_id='zork1.z5', - seed=0, - debug_mode=True - ) - else: - self.priorzero_cfg, self.priorzero_create_cfg = get_priorzero_config( - env_id='zork1.z5', - seed=0, - enable_llm=True, - enable_rft=True, - debug_mode=False - ) - - # Hybrid-specific settings - self.wm_training_mode = "parallel" - self.wm_train_freq = 1 - self.llm_train_freq = 5 - self.use_orz_trainer = ORZ_AVAILABLE - - # vLLM settings - self.use_vllm = VLLM_AVAILABLE - self.vllm_required = False # Set to True if vLLM is required - - # ORZ-specific settings (only used if ORZ_AVAILABLE) - if ORZ_AVAILABLE: - self.orz_rollout_batch_size = 32 if DEBUG_MODE else 128 - self.orz_train_batch_size = 8 if DEBUG_MODE else 32 - self.orz_actor_lr = 1e-6 - self.orz_critic_lr = 5e-6 - self.orz_num_episodes = 2 if DEBUG_MODE else 10 - - -# ============================================================================== -# ORZ Data Adapter and Dataset -# ============================================================================== - -class GameSegmentToORZAdapter: - """ - Convert PriorZero game_segments to ORZ-compatible format. - """ - - @staticmethod - def convert_segments_to_prompts(game_segments: List[Any], tokenizer) -> List[Dict]: - """ - Convert game_segments to ORZ prompt format. - - Args: - game_segments: List of GameSegment from PriorZero - tokenizer: HuggingFace tokenizer - - Returns: - List of ORZ-compatible prompt dictionaries - """ - prompts = [] - - for segment in game_segments: - # Extract raw observations if available - if hasattr(segment, 'raw_obs_segment') and segment.raw_obs_segment: - for i, (obs, action) in enumerate(zip( - segment.raw_obs_segment, - segment.action_segment - )): - # Create ORZ format prompt - prompt_dict = { - "prompt": [{"value": obs}], - "final_answer": action, - "file_name": f"segment_{id(segment)}_step_{i}" - } - prompts.append(prompt_dict) - - return prompts - - @staticmethod - def extract_training_data(game_segments: List[Any]) -> Dict[str, List]: - """ - Extract training data from game_segments for ORZ. - - Returns: - Dictionary containing: - - states: List of state descriptions - - actions: List of actions taken - - rewards: List of rewards received - - mcts_policies: List of MCTS visit distributions - """ - training_data = { - 'states': [], - 'actions': [], - 'rewards': [], - 'mcts_policies': [] - } - - for segment in game_segments: - # Extract raw observations (states) - if hasattr(segment, 'raw_obs_segment'): - training_data['states'].extend(segment.raw_obs_segment) - - # Extract actions - if hasattr(segment, 'action_segment'): - training_data['actions'].extend(segment.action_segment) - - # Extract rewards - if hasattr(segment, 'reward_segment'): - training_data['rewards'].extend(segment.reward_segment) - - # Extract MCTS policies - if hasattr(segment, 'mcts_policy_segment'): - training_data['mcts_policies'].extend(segment.mcts_policy_segment) - - return training_data - - -# Only define dataset classes if ORZ is available -if ORZ_AVAILABLE: - from jinja2 import Template - - class JerichoPromptDataset(PromptDataset): - """ - Custom dataset for Jericho text adventure games in ORZ format. - Adapts PriorZero game_segments to ORZ PPO training format. - """ - - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) - - def process_dialogue(self, dialogue: dict): - """ - Process a single dialogue (observation + action pair) into ORZ format. - - Args: - dialogue: Dict with 'prompt', 'final_answer', 'file_name' - - Returns: - prompt: Formatted prompt string - extra: Dict with answer and metadata - """ - # Template for Jericho text adventure prompts - prompt_template_jinja = """\ -{{bos_token}}A conversation between User and Assistant. The User is playing a text adventure game \ -and needs to decide the next action. The Assistant carefully analyzes the current game state, \ -considers the available actions, and recommends the best action to take. \ -The reasoning process is enclosed within tags, and the recommended action \ -is enclosed within tags. For example: \ - The player is in a dark room and needs light. The lamp is available. \ - take lamp . User: {{prompt}} -Assistant: \ -""" - - prompt_instruction_template_jinja = """\ -Current game state: -{{prompt}} - -What is the best action to take? Put your answer inside tags. -""" - - # Validate dialogue format - assert isinstance(dialogue, dict), "dialogue must be a dict" - assert "prompt" in dialogue, "dialogue must contain prompt" - assert "final_answer" in dialogue, "dialogue must contain final_answer" - - # Build prompt - prompt_instruction_template = Template(prompt_instruction_template_jinja) - prompt_instruction = prompt_instruction_template.render( - prompt=dialogue["prompt"][0]["value"] - ) - - prompt_template = Template(prompt_template_jinja) - if self.tokenizer.bos_token_id is None: - bos_token = "" - else: - bos_token = self.tokenizer.decode([self.tokenizer.bos_token_id]) - - prompt = prompt_template.render( - bos_token=bos_token, - prompt=prompt_instruction - ) - - extra = { - "answer": dialogue["final_answer"], - "file_name": dialogue.get("file_name", "unknown") - } - - return prompt, extra - - -# ============================================================================== -# Main Training Function -# ============================================================================== - -async def train_priorzero_orz_complete( - cfg: dict, - create_cfg: dict, - hybrid_cfg: HybridTrainingConfig, - seed: int = 0, - max_train_iter: int = 10000, - max_env_step: Optional[int] = int(1e10), - enable_save: bool = True, -): - """ - Main hybrid training function with complete ORZ integration. - """ - # ================================================================== - # 1. Compile Configuration - # ================================================================== - cfg = compile_config(cfg, seed=seed, auto=True, create_cfg=create_cfg) - - # ================================================================== - # 2. Create vLLM Engine (optional) - Based on priorzero_entry.py - # ================================================================== - vllm_engine = None - - if hybrid_cfg.use_vllm and VLLM_AVAILABLE: - logger.info("Creating vLLM engine...") - - # [ROBUST FIX] Handle shared GPU environment - # Solution: Use alternative initialization method with fallback - tensor_parallel = cfg.policy.llm_policy_cfg.vllm_tensor_parallel_size - distributed_backend = "ray" if tensor_parallel > 1 else None - - # [ROBUST FIX] Lower GPU memory utilization in shared environment - gpu_mem_util = cfg.policy.llm_policy_cfg.gpu_memory_utilization - if gpu_mem_util > 0.85: - gpu_mem_util = 0.75 # More conservative - logger.info(f"✓ Adjusted GPU memory utilization to {gpu_mem_util} for stability") - - # [ROBUST FIX] Use vLLM V0 engine for stability (as in priorzero_entry.py) - use_v1_env = os.environ.get('VLLM_USE_V1', None) - if use_v1_env is None: - # Only set if not already set by user - os.environ['VLLM_USE_V1'] = '0' - logger.info("✓ Using vLLM V0 engine for stability") - - # Fix tokenizers parallelism warning - os.environ['TOKENIZERS_PARALLELISM'] = 'false' - - try: - from vllm.engine.arg_utils import AsyncEngineArgs - - engine_args = AsyncEngineArgs( - model=cfg.policy.llm_policy_cfg.pretrain_llm_path, - tensor_parallel_size=tensor_parallel, - gpu_memory_utilization=gpu_mem_util, - distributed_executor_backend=distributed_backend, - trust_remote_code=True, - enable_prefix_caching=False, - enforce_eager=False, - ) - vllm_engine = AsyncLLMEngine.from_engine_args(engine_args) - logger.info(f"✓ vLLM Engine created (backend: {distributed_backend or 'default'})") - - except (ValueError, RuntimeError) as e: - if "VLLM_USE_V1" in str(e) or "memory profiling" in str(e): - # Fallback: Try without V1 env var or with eager mode - logger.warning(f"⚠️ Initial vLLM initialization failed: {e}") - logger.info("Retrying with alternative configuration...") - - if 'VLLM_USE_V1' in os.environ: - del os.environ['VLLM_USE_V1'] - - try: - engine_args = AsyncEngineArgs( - model=cfg.policy.llm_policy_cfg.pretrain_llm_path, - tensor_parallel_size=tensor_parallel, - gpu_memory_utilization=gpu_mem_util * 0.9, # Even more conservative - distributed_executor_backend=distributed_backend, - trust_remote_code=True, - enable_prefix_caching=False, - enforce_eager=True, # Force eager mode as fallback - ) - vllm_engine = AsyncLLMEngine.from_engine_args(engine_args) - logger.info(f"✓ vLLM Engine created with fallback configuration") - except Exception as e2: - logger.error(f"❌ Failed to create vLLM engine with fallback: {e2}") - if hybrid_cfg.vllm_required: - raise - logger.warning("Continuing without vLLM (LLM prior will be disabled)") - else: - logger.error(f"❌ Failed to create vLLM engine: {e}") - import traceback - logger.error(f"Full traceback:\n{traceback.format_exc()}") - if hybrid_cfg.vllm_required: - raise - logger.warning("Continuing without vLLM (LLM prior will be disabled)") - else: - logger.info("vLLM disabled or not available - continuing without LLM inference") - - # ================================================================== - # 3. Create Environments - # ================================================================== - logger.info("Creating environments...") - env_fn, collector_env_cfg, evaluator_env_cfg = get_vec_env_setting(cfg.env) - - collector_env = create_env_manager( - cfg.env.manager, - [partial(env_fn, cfg=c) for c in collector_env_cfg] - ) - evaluator_env = create_env_manager( - cfg.env.manager, - [partial(env_fn, cfg=c) for c in evaluator_env_cfg] - ) - - # Seed environments - collector_env.seed(seed) - evaluator_env.seed(seed, dynamic_seed=False) - set_pkg_seed(seed, use_cuda=True) - logger.info(f"✓ Environments created and seeded (seed={seed})") - - # ================================================================== - # 4. Create Policy, Buffer, and Components - # ================================================================== - logger.info("Creating policy, buffer, and components...") - - # Create policy - policy = create_policy( - cfg.policy, - enable_field=['learn', 'collect', 'eval'] - ) - logger.info("✓ Policy created") - - # Create TensorBoard logger - os.makedirs(f'./{cfg.exp_name}/log/', exist_ok=True) - tb_logger = SummaryWriter( - os.path.join(f'./{cfg.exp_name}/log/', 'serial') - ) if get_rank() == 0 else None - logger.info(f"✓ TensorBoard logger: ./{cfg.exp_name}/log/") - - # Create learner (for world model training) - learner = BaseLearner( - cfg.policy.learn.learner, - policy.learn_mode, - tb_logger, - exp_name=cfg.exp_name - ) - logger.info("✓ BaseLearner created") - - # Create replay buffer - replay_buffer = PriorZeroGameBufferOptimized(cfg.policy) - logger.info("✓ PriorZero replay buffer created") - - # Create collector - collector = PriorZeroCollector( - env=collector_env, - policy=policy.collect_mode, - tb_logger=tb_logger, - exp_name=cfg.exp_name, - vllm_engine=vllm_engine, # May be None - policy_config=cfg.policy, - debug_mode=cfg.get('debug_mode', False), - ) - logger.info("✓ Collector created") - - # Create evaluator - evaluator = PriorZeroEvaluator( - eval_freq=cfg.policy.eval_freq, - n_evaluator_episode=cfg.env.n_evaluator_episode, - stop_value=cfg.env.stop_value, - env=evaluator_env, - policy=policy.eval_mode, - tb_logger=tb_logger, - exp_name=cfg.exp_name, - vllm_engine=vllm_engine, # May be None - ) - logger.info("✓ Evaluator created") - - # Call learner's before_run hook - learner.call_hook('before_run') - - # ================================================================== - # 5. Initialize ORZ Trainer (if available) - # ================================================================== - orz_trainer = None - orz_adapter = GameSegmentToORZAdapter() - orz_tokenizer = None - orz_strategy = None - - if hybrid_cfg.use_orz_trainer and ORZ_AVAILABLE: - logger.info("="*80) - logger.info("Initializing ORZ RayPPOTrainer for LLM training...") - logger.info("="*80) - - try: - # Initialize Ray if not already running - if not ray.is_initialized(): - ray.init(ignore_reinit_error=True) - logger.info("✓ Ray initialized") - - # Create ORZ tokenizer - orz_tokenizer = AutoTokenizer.from_pretrained( - cfg.policy.llm_policy_cfg.pretrain_llm_path, - trust_remote_code=True - ) - if orz_tokenizer.pad_token is None: - orz_tokenizer.pad_token = orz_tokenizer.eos_token - logger.info("✓ ORZ tokenizer created") - - # Create ORZ strategy (DeepSpeed config) - from orz.ppo.utils import get_strategy - orz_strategy = get_strategy({ - 'zero_stage': 2, - 'bf16': True, - 'gradient_checkpointing': True, - }) - logger.info("✓ ORZ strategy created") - - # Create ORZ configuration (matching ORZ's PPOExpConfig pattern) - from dataclasses import dataclass, field - from omegaconf.listconfig import ListConfig - - @dataclass - class ORZConfig: - """Simplified ORZ config for PriorZero integration""" - # Resource settings (simplified for single-node) - total_num_nodes: int = 1 - ref_num_nodes: int = 1 - ref_num_gpus_per_node: int = 1 - actor_num_nodes: int = 1 - actor_num_gpus_per_node: int = 1 - critic_num_nodes: int = 1 - critic_num_gpus_per_node: int = 1 - colocate_all: bool = True - colocate_critic_reward: bool = True - colocate_actor_ref: bool = True - vllm_num_engines: int = 1 - vllm_tensor_parallel_size: int = 1 - zero_stage: int = 2 - adam_offload: bool = False - - # Model paths - pretrain: str = cfg.policy.llm_policy_cfg.pretrain_llm_path - reward_pretrain: Optional[str] = None - critic_pretrain: Optional[str] = cfg.policy.llm_policy_cfg.pretrain_llm_path - - # Save/log paths - save_interval: int = 50 - ckpt_path: str = f'./{cfg.exp_name}/orz_ckpt' - save_path: str = f'./{cfg.exp_name}/orz_save' - tensorboard_log_dir: str = f'./{cfg.exp_name}/orz_log' - - # Training settings - actor_learning_rate: float = hybrid_cfg.orz_actor_lr if hasattr(hybrid_cfg, 'orz_actor_lr') else 1e-6 - critic_learning_rate: float = hybrid_cfg.orz_critic_lr if hasattr(hybrid_cfg, 'orz_critic_lr') else 5e-6 - num_warmup_steps: int = 50 - prompt_max_len: int = 2048 - enable_prefix_caching: bool = False - update_ref_every_epoch: bool = True - advantage_normalize: bool = True - - # Episode settings - num_episodes: int = hybrid_cfg.orz_num_episodes if hasattr(hybrid_cfg, 'orz_num_episodes') else 2 - rollout_batch_size: int = hybrid_cfg.orz_rollout_batch_size if hasattr(hybrid_cfg, 'orz_rollout_batch_size') else 32 - n_samples_per_prompt: int = 8 if DEBUG_MODE else 32 - micro_rollout_batch_size: int = 2 - policy_update_steps: int = 1 - critic_update_steps: int = 1 if DEBUG_MODE else 12 - micro_train_batch_size: int = 1 - micro_forward_batch_size: int = 1 - freezing_actor_steps: int = -1 - - # KL settings - init_kl_coef: float = 0 - kl_loss_coef: float = 0.0 - use_kl_loss: bool = False - use_kl_estimator_k3: bool = True - - # Eval settings - enable_eval: bool = False # Disable ORZ eval (use PriorZero's) - eval_interval: int = 100 - - # Generation settings - packing_max_len: int = 8192 - generate_max_len: int = cfg.policy.llm_policy_cfg.generate_max_len - max_len: int = 4096 - temperature: float = 1.0 - top_p: float = 1.0 - top_k: int = -1 - stop: ListConfig = field(default_factory=lambda: ListConfig([""])) - - # GRPO settings - use_grpo: bool = False - gamma: float = 1.0 - lambd: float = 1.0 - - # vLLM settings - gpu_memory_utilization: float = 0.3 - - # Custom settings for compute_reward_fn - use_compute_reward_fn: bool = True - use_orm_score: bool = False - - orz_cfg = ORZConfig() - - # Create directories for ORZ - os.makedirs(orz_cfg.ckpt_path, exist_ok=True) - os.makedirs(orz_cfg.save_path, exist_ok=True) - os.makedirs(orz_cfg.tensorboard_log_dir, exist_ok=True) - - logger.info("✓ ORZ config created") - logger.info(f" - Model: {orz_cfg.pretrain}") - logger.info(f" - Rollout batch: {orz_cfg.rollout_batch_size}") - logger.info(f" - Episodes: {orz_cfg.num_episodes}") - - # Note: Full RayPPOTrainer initialization requires: - # 1. Creating vLLM engines for distributed inference - # 2. Creating initial dataset from game_segments - # 3. Initializing Ray actors (will be done lazily on first training call) - # - # We defer full initialization until we have actual game_segments to train on - logger.info("✓ ORZ trainer components ready") - logger.info(" (Full RayPPOTrainer will be initialized on first training iteration)") - - except Exception as e: - logger.error(f"❌ ORZ trainer initialization failed: {e}") - import traceback - logger.error(traceback.format_exc()) - logger.warning("Falling back to PriorZero's built-in LLM training") - hybrid_cfg.use_orz_trainer = False - - # ================================================================== - # 6. Main Training Loop - # ================================================================== - logger.info("="*80) - logger.info("Starting PriorZero-ORZ Complete Training") - logger.info("="*80) - logger.info(f"Experiment: {cfg.exp_name}") - logger.info(f"Max iterations: {max_train_iter}") - logger.info(f"Training mode: {hybrid_cfg.wm_training_mode}") - logger.info(f"Use ORZ trainer: {hybrid_cfg.use_orz_trainer}") - logger.info(f"Use vLLM: {vllm_engine is not None}") - logger.info(f"LLM model: {cfg.policy.llm_policy_cfg.pretrain_llm_path}") - logger.info(f"World model: UniZero") - logger.info("="*80) - - # Training state - best_eval_reward = -float('inf') - total_game_segments_collected = 0 - - try: - while learner.train_iter < max_train_iter and collector.envstep < max_env_step: - current_iter = learner.train_iter - - # ============================================================== - # Step 1: Evaluation (if needed) - # ============================================================== - if current_iter > 0 and evaluator.should_eval(current_iter): - logger.info(f"\n{'='*60}") - logger.info(f"[Iter {current_iter}] Evaluating...") - logger.info(f"{'='*60}") - - eval_result = await evaluator.eval( - save_ckpt_fn=learner.save_checkpoint if enable_save else None, - train_iter=current_iter, - envstep=collector.envstep - ) - - if eval_result is not None: - stop, eval_reward_dict = eval_result - mean_reward = eval_reward_dict.get('reward_mean', 0) - logger.info(f"✓ Evaluation: reward_mean={mean_reward:.2f}") - - if mean_reward > best_eval_reward: - best_eval_reward = mean_reward - logger.info(f"🎯 New best reward: {best_eval_reward:.2f}") - - if stop: - logger.info(f"🎉 Training converged! (reward >= {cfg.env.stop_value})") - break - - # ============================================================== - # Step 2: Collect Data using MCTS - # ============================================================== - logger.info(f"\n[Iter {current_iter}] Collecting data...") - - collect_kwargs = { - 'temperature': 0.25, - 'epsilon': 0.0 - } - - try: - new_data = await collector.collect( - train_iter=current_iter, - policy_kwargs=collect_kwargs - ) - except Exception as e: - logger.error(f"❌ Collection failed: {e}") - logger.warning("Skipping this iteration...") - continue - - # Add to replay buffer - from lzero.entry.utils import calculate_update_per_collect - update_per_collect = calculate_update_per_collect(cfg, new_data, world_size=1) - - # Update buffer - replay_buffer.push_game_segments(new_data) - logger.info( - f"✓ Collected {len(new_data)} segments " - f"(total: {replay_buffer.get_num_of_game_segments()} segments, " - f"{replay_buffer.get_num_of_transitions()} transitions)" - ) - - total_game_segments_collected += len(new_data) - - # ============================================================== - # Step 3: World Model Training - # ============================================================== - if current_iter % hybrid_cfg.wm_train_freq == 0: - if replay_buffer.get_num_of_transitions() >= cfg.policy.batch_size: - logger.info(f"[Iter {current_iter}] Training world model...") - - # Sample and train - for _ in range(update_per_collect): - train_data = replay_buffer.sample( - cfg.policy.batch_size, - policy - ) - - # Train (includes both WM and LLM in PriorZero) - log_dict = learner.train(train_data, collector.envstep) - - # Log to TensorBoard - if tb_logger and get_rank() == 0: - for k, v in log_dict.items(): - tb_logger.add_scalar(f'train/{k}', v, collector.envstep) - - logger.info( - f"✓ WM training done - " - f"wm_loss: {log_dict.get('wm_total_loss', 0):.4f}, " - f"llm_sft_loss: {log_dict.get('llm_sft_loss', 0):.4f}" - ) - else: - logger.info(f"Skipping training - not enough data yet") - - # ============================================================== - # Step 4: LLM Training with ORZ (if enabled) - # ============================================================== - if (hybrid_cfg.use_orz_trainer and orz_trainer is not None and - current_iter % hybrid_cfg.llm_train_freq == 0 and - current_iter > 0): - logger.info(f"[Iter {current_iter}] Training LLM with ORZ...") - - try: - # Extract game_segments from recent collections - training_data = orz_adapter.extract_training_data(new_data) - num_samples = len(training_data['states']) - - if num_samples > 0: - logger.info(f" Extracted {num_samples} training samples for ORZ") - - # Initialize ORZ trainer on first use (lazy initialization) - if orz_trainer is None: - logger.info(" Initializing ORZ RayPPOTrainer...") - - # Convert game_segments to ORZ dataset format - dialogues = orz_adapter.convert_segments_to_prompts( - new_data, - orz_tokenizer - ) - - # Create ORZ dataset - orz_dataset = JerichoPromptDataset( - dialogues, - orz_tokenizer, - orz_cfg.prompt_max_len, - orz_strategy, - pretrain_mode=False, - num_processors=1 - ) - - # Create custom reward trainer - from orz.exps.examples.ppo.ppo_base_exp import BasePPOExp - - class JerichoRewardTrainer(RayPPOTrainer): - """Custom reward trainer for Jericho text adventures""" - - async def custom_reward_fn( - self, - prompts: List[str], - outputs: List[Any], - extras: List[dict], - reward_model_fn, - ): - """ - Compute rewards for Jericho actions. - Reward is 1.0 if action matches ground truth, else 0.0 - """ - import torch - scores = [] - responses = [] - - for output, extra in zip(outputs, extras): - response = output["response"] - responses.append(response) - - # Extract action from response - # Look for ... tags - import re - pattern = re.compile(r"(.*?)", re.DOTALL) - matches = re.findall(pattern, response) - predicted_action = matches[-1].strip() if matches else "" - - # Ground truth action - true_action = extra["answer"] - - # Simple exact match for now - # TODO: Could use fuzzy matching or LLM-based similarity - score = 1.0 if predicted_action.lower() == true_action.lower() else 0.0 - scores.append(score) - - # Log statistics - avg_score = sum(scores) / len(scores) if scores else 0.0 - logger.info(f" ORZ reward - avg: {avg_score:.3f}, samples: {len(scores)}") - - # Create score tensors (reward only on last token) - output_tokens = self._tokenize(responses, self.cfg.generate_max_len, padding=False)["input_ids"] - score_tensors = [] - for score, output_token in zip(scores, output_tokens): - score_tensor = torch.zeros(len(output_token)) - if len(output_token) > 0: - score_tensor[-1] = score - score_tensors.append(score_tensor) - - # Remove empty responses - res_prompts, res_responses, res_score_tensors = [], [], [] - for prompt, response, score_tensor in zip(prompts, responses, score_tensors): - if len(response) > 0: - res_prompts.append(prompt) - res_responses.append(response) - res_score_tensors.append(score_tensor) - - return res_prompts, res_responses, res_score_tensors - - # Create vLLM engines for ORZ - logger.info(" Creating vLLM inference engines for ORZ...") - from orz.exps.examples.ppo.ppo_base_exp import BasePPOExp - - # Use BasePPOExp helper to create engines - class TempExp(BasePPOExp): - def __init__(self): - self.cfg = orz_cfg - self.tokenizer = orz_tokenizer - self.strategy = orz_strategy - - temp_exp = TempExp() - vllm_engines = temp_exp.create_inference_engine() - logger.info(f" ✓ Created {len(vllm_engines)} vLLM engines") - - # Get colocate placement groups if needed - colocate_pg = temp_exp.get_colocate_pg if orz_cfg.colocate_all else None - - # Create ORZ trainer - orz_trainer = JerichoRewardTrainer( - cfg=orz_cfg, - strategy=orz_strategy, - tokenizer=orz_tokenizer, - train_dataset=orz_dataset, - eval_dataset=None, # No separate eval for now - vllm_engines=vllm_engines, - colocate_pg=colocate_pg - ) - - logger.info(" ✓ ORZ RayPPOTrainer initialized") - - # Run ORZ training for one episode - logger.info(f" Running ORZ PPO training (episode {current_iter // hybrid_cfg.llm_train_freq})...") - - # Train using ORZ's fit_episode method - # Note: This will do full PPO update with actor/critic training - await orz_trainer.fit_episode() - - logger.info(f" ✓ ORZ training completed for iteration {current_iter}") - - else: - logger.warning(" No training samples extracted from game_segments") - - except Exception as e: - logger.error(f" ✗ ORZ training failed: {e}") - import traceback - logger.error(traceback.format_exc()) - logger.warning(" Continuing with PriorZero LLM training only") - - # ============================================================== - # Step 5: Logging and Checkpointing - # ============================================================== - if current_iter % 10 == 0: - logger.info(f"\n{'='*60}") - logger.info(f"Progress Summary (Iter {current_iter})") - logger.info(f"{'='*60}") - logger.info(f"Env steps: {collector.envstep}") - logger.info(f"Game segments collected: {total_game_segments_collected}") - logger.info(f"Buffer size: {replay_buffer.get_num_of_transitions()} transitions") - logger.info(f"Best eval reward: {best_eval_reward:.2f}") - logger.info(f"{'='*60}\n") - - # Save checkpoint periodically - if enable_save and current_iter % 100 == 0 and current_iter > 0: - logger.info(f"[Iter {current_iter}] Saving checkpoint...") - learner.save_checkpoint(collector.envstep) - logger.info("✓ Checkpoint saved") - - except KeyboardInterrupt: - logger.info("\n⚠️ Training interrupted by user") - except Exception as e: - logger.error(f"\n❌ Training failed with error: {e}") - import traceback - traceback.print_exc() - raise - finally: - # ============================================================== - # Cleanup - # ============================================================== - logger.info("\nCleaning up...") - - # Save final checkpoint - if enable_save: - logger.info("Saving final checkpoint...") - try: - learner.save_checkpoint(collector.envstep) - except Exception as e: - logger.error(f"Failed to save checkpoint: {e}") - - # Close environments - try: - collector_env.close() - evaluator_env.close() - except Exception as e: - logger.error(f"Failed to close environments: {e}") - - # Close loggers - if tb_logger: - try: - tb_logger.close() - except Exception as e: - logger.error(f"Failed to close tensorboard: {e}") - - logger.info("✓ Cleanup complete") - logger.info("="*80) - logger.info("Training finished!") - logger.info(f"Total iterations: {learner.train_iter}") - logger.info(f"Total env steps: {collector.envstep}") - logger.info(f"Best eval reward: {best_eval_reward:.2f}") - logger.info("="*80) - - -# ============================================================================== -# Entry Point -# ============================================================================== - -async def main(): - """Main entry point.""" - # Create hybrid configuration - hybrid_cfg = HybridTrainingConfig() - - # Run training - await train_priorzero_orz_complete( - cfg=hybrid_cfg.priorzero_cfg, - create_cfg=hybrid_cfg.priorzero_create_cfg, - hybrid_cfg=hybrid_cfg, - seed=0, - max_train_iter=10000 if not DEBUG_MODE else 100, - enable_save=True, - ) - - -if __name__ == "__main__": - logger.info("="*80) - logger.info("PriorZero-ORZ Complete Training Pipeline") - logger.info("="*80) - logger.info(f"Debug mode: {DEBUG_MODE}") - logger.info(f"ORZ available: {ORZ_AVAILABLE}") - logger.info(f"vLLM available: {VLLM_AVAILABLE}") - logger.info("="*80) - - # Run async training - asyncio.run(main()) diff --git a/zoo/jericho/priorzero/priorzero_policy.py b/zoo/jericho/priorzero/priorzero_policy.py deleted file mode 100644 index 26e50e060..000000000 --- a/zoo/jericho/priorzero/priorzero_policy.py +++ /dev/null @@ -1,1466 +0,0 @@ -# priorzero_policy.py -""" -[PRIORZERO] PriorZero Policy Implementation - -This module implements the PriorZero policy that combines: -1. UniZero world model for planning in latent space -2. LLM policy model for providing high-quality action priors - -Key Features: -- Dual-model training: world model + LLM policy -- LLM-guided MCTS: inject LLM priors into MCTS root node -- SFT + RFT: supervised fine-tuning with MCTS policies + reinforcement fine-tuning with environment rewards -- Full alignment with UniZero implementation - -Author: PriorZero Team -Date: 2025-01-20 -""" - -import copy -import re -import sys -import logging -from pathlib import Path -from typing import List, Dict, Any, Tuple, Union, Optional - -# [CRITICAL] Ensure local LightZero is used -from ensure_local_lightzero import ensure_local_lightzero -ensure_local_lightzero() - -import numpy as np -import torch -import torch.nn.functional as F -from ding.utils import POLICY_REGISTRY -from ding.model import model_wrap -from transformers import AutoTokenizer, AutoModelForCausalLM -from peft import get_peft_model, LoraConfig, TaskType - -# Import from local LightZero -from lzero.policy.unizero import UniZeroPolicy as OriginalUniZeroPolicy -from lzero.policy import ( - phi_transform, - InverseScalarTransform, - scalar_transform, # [PRIORZERO] Added for reward/value transformation - DiscreteSupport, # [PRIORZERO] Added for categorical distribution support - to_torch_float_tensor, - mz_network_output_unpack -) -from lzero.policy.utils import select_action -from lzero.mcts import UniZeroMCTSCtree as MCTSCtree -from lzero.entry.utils import initialize_zeros_batch -# Import UniZeroModel to ensure it's registered in MODEL_REGISTRY -import lzero.model.unizero_model # noqa: F401 - - -# ============================================================================== -# Helper Functions for LLM Prior Processing -# ============================================================================== - -def parse_llm_action_ranking( - text: str, - action_map: Dict[str, int], - action_space_size: int, - fallback_to_uniform: bool = True -) -> np.ndarray: - """ - [PRIORZERO-NEW] - Parse LLM generated action ranking text into a policy distribution. - - Args: - text: LLM generated text with ranked actions (e.g., "1. take key\\n2. go north") - action_map: Mapping from action text to action index - action_space_size: Size of the action space - fallback_to_uniform: If True, return uniform distribution when no valid action found - - Returns: - policy: Probability distribution over actions (shape: [action_space_size]) - """ - # Extract ranked actions using regex - # Supports formats: "1. action", "1) action", "1: action" - ranked_actions = re.findall(r'(?:^|\n)\s*\d+[\.\):\s]+(.+?)(?=\n|$)', text, re.MULTILINE) - - policy = np.zeros(action_space_size, dtype=np.float32) - found_count = 0 - - for rank, action_text in enumerate(ranked_actions): - action_text = action_text.strip().lower() - - # Try exact match first - if action_text in action_map: - action_idx = action_map[action_text] - # Assign decreasing weights (higher rank = higher weight) - policy[action_idx] = len(ranked_actions) - rank - found_count += 1 - else: - # Try fuzzy matching (find best substring match) - best_match_score = 0 - best_action_idx = None - for candidate_text, candidate_idx in action_map.items(): - if candidate_text in action_text or action_text in candidate_text: - score = len(set(candidate_text.split()) & set(action_text.split())) - if score > best_match_score: - best_match_score = score - best_action_idx = candidate_idx - - if best_action_idx is not None: - policy[best_action_idx] = len(ranked_actions) - rank - found_count += 1 - - # Normalize to probability distribution - if policy.sum() > 0: - policy /= policy.sum() - elif fallback_to_uniform: - # If LLM didn't generate any valid actions, return uniform distribution - policy = np.ones(action_space_size, dtype=np.float32) / action_space_size - - return policy - - -def format_mcts_policy_to_text( - mcts_policy: np.ndarray, - action_inv_map: Dict[int, str], - top_k: int = 5 -) -> str: - """ - [PRIORZERO-NEW] - Convert MCTS policy vector into ranked action text for SFT training. - - Args: - mcts_policy: MCTS visit count distribution (shape: [action_space_size]) - action_inv_map: Mapping from action index to action text - top_k: Number of top actions to include - - Returns: - Formatted text with ranked actions (e.g., "1. take key\\n2. go north\\n...") - """ - # Sort actions by policy probability (descending) - sorted_indices = np.argsort(mcts_policy)[::-1] - - output_lines = [] - rank = 1 - for idx in sorted_indices: - if mcts_policy[idx] > 0 and rank <= top_k: - action_text = action_inv_map.get(idx, f"action_{idx}") - output_lines.append(f"{rank}. {action_text}") - rank += 1 - - return "\n".join(output_lines) if output_lines else "No valid actions found." - - -def build_llm_prompt( - current_obs: str, - history: Optional[List[Tuple[str, str, float]]] = None, - action_descriptions: Optional[Dict[str, str]] = None, - use_cot: bool = True -) -> str: - """ - [PRIORZERO-NEW] - Build a high-quality prompt for LLM to generate action ranking. - - Args: - current_obs: Current observation text - history: List of (observation, action, reward) tuples - action_descriptions: Optional descriptions for each action - use_cot: Whether to encourage chain-of-thought reasoning - - Returns: - Formatted prompt string - """ - prompt_parts = [] - - # System instruction - prompt_parts.append( - "You are an expert player in a text-based adventure game. " - "Your goal is to maximize the score by taking the best actions." - ) - - # Add history if available - if history and len(history) > 0: - prompt_parts.append("\n=== Recent History ===") - for i, (obs, action, reward) in enumerate(history[-5:]): # Last 5 steps - prompt_parts.append(f"Step {i+1}:") - prompt_parts.append(f" Observation: {obs[:100]}...") # Truncate long obs - prompt_parts.append(f" Action: {action}") - prompt_parts.append(f" Reward: {reward}") - - # Current observation - prompt_parts.append("\n=== Current Situation ===") - prompt_parts.append(current_obs) - - # Task instruction - if use_cot: - prompt_parts.append( - "\n=== Task ===\n" - "Think step-by-step:\n" - "1. Analyze the current situation and your goal\n" - "2. Consider what actions might help you progress\n" - "3. Rank the best actions in order of priority\n" - "\nProvide your analysis and then list the top 5 actions in this format:\n" - "1. [first action]\n" - "2. [second action]\n" - "..." - ) - else: - prompt_parts.append( - "\n=== Task ===\n" - "List the top 5 best actions in order of priority:\n" - "1. [first action]\n" - "2. [second action]\n" - "..." - ) - - return "\n".join(prompt_parts) - - -# ============================================================================== -# PriorZero Policy Class -# ============================================================================== - -@POLICY_REGISTRY.register('priorzero', force_overwrite=True) -class PriorZeroPolicy(OriginalUniZeroPolicy): - """ - [PRIORZERO-MODIFIED] - PriorZero policy that combines UniZero world model with LLM policy. - - Architecture: - - UniZero World Model: Learns latent dynamics, value, and policy in latent space - - LLM Policy Model: Provides high-quality action priors based on language understanding - - Training: - - World Model: Trained with standard UniZero losses (value, policy, reward, latent) - - LLM: Trained with SFT (using MCTS policies) + RFT (using environment rewards) - - Inference: - - LLM generates action ranking → converted to policy prior - - Policy prior injected into MCTS root node - - MCTS search refines the policy → selects best action - """ - - config = dict( - **OriginalUniZeroPolicy.config, - # LLM-specific config - llm_policy_cfg=dict( - pretrain_llm_path="Qwen/Qwen1.5-1.8B-Chat", - use_lora=False, # Whether to use LoRA for efficient fine-tuning - lora_r=8, - lora_alpha=16, - lora_dropout=0.05, - llm_learning_rate=1e-6, - llm_weight_decay=0.01, - llm_loss_weight=0.5, # Weight of LLM loss in total loss - rft_loss_weight=0.3, # Weight of RFT loss in total loss - prompt_max_len=2048, - generate_max_len=128, - history_length=5, # Number of recent steps to include in prompt - use_cot=True, # Whether to use chain-of-thought prompting - sft_target='mcts_policy', # 'mcts_policy' or 'oracle_policy' - enable_rft=True, # Whether to enable RFT training - ), - ) - - def __init__(self, cfg: Dict, model: torch.nn.Module = None, enable_field: List[str] = None): - # [PRIORZERO-NEW] Initialize LLM-related attributes BEFORE super().__init__ - # because super().__init__ will call _init_learn which needs these attributes - self.llm_policy_model = None - self.llm_tokenizer = None - self._optimizer_llm = None - self._lr_scheduler_llm = None - self.llm_policy_cfg = cfg.llm_policy_cfg # Set from cfg, not self._cfg (not set yet) - - # Action mapping (will be set from config) - self.action_map = None # str -> int - self.action_inv_map = None # int -> str - - # Call parent init (this will trigger _init_learn, _init_collect, _init_eval) - super().__init__(cfg, model, enable_field) - - def _init_learn(self) -> None: - """ - [PRIORZERO-MODIFIED] - Initialize both UniZero world model and LLM policy model with their optimizers. - Align with UniZero implementation - use logging instead of self._logger. - """ - import logging - - # ====================================================================== - # 1. Initialize UniZero World Model (from parent class) - # ====================================================================== - super()._init_learn() - logging.info("✓ UniZero World Model and optimizer initialized") - - # [PRIORZERO-FIX] Ensure scalar transform handles are initialized - # These are normally initialized in UniZeroPolicy.__init__ but we need to ensure they exist - if not hasattr(self, 'value_support') or self.value_support is None: - self.value_support = DiscreteSupport(*self._cfg.model.value_support_range, self._cfg.device) - if not hasattr(self, 'reward_support') or self.reward_support is None: - self.reward_support = DiscreteSupport(*self._cfg.model.reward_support_range, self._cfg.device) - if not hasattr(self, 'value_inverse_scalar_transform_handle'): - self.value_inverse_scalar_transform_handle = InverseScalarTransform( - self.value_support, self._cfg.model.categorical_distribution - ) - if not hasattr(self, 'reward_inverse_scalar_transform_handle'): - self.reward_inverse_scalar_transform_handle = InverseScalarTransform( - self.reward_support, self._cfg.model.categorical_distribution - ) - logging.info("✓ Scalar transform handles verified/initialized") - - # ====================================================================== - # 2. [PRIORZERO-NEW] Initialize LLM Policy Model - # ====================================================================== - logging.info(f"Loading LLM from: {self.llm_policy_cfg.pretrain_llm_path}") - - # Load tokenizer - self.llm_tokenizer = AutoTokenizer.from_pretrained( - self.llm_policy_cfg.pretrain_llm_path, - trust_remote_code=True, - padding_side='left' # For batch generation - ) - if self.llm_tokenizer.pad_token is None: - self.llm_tokenizer.pad_token = self.llm_tokenizer.eos_token - - # Load LLM - self.llm_policy_model = AutoModelForCausalLM.from_pretrained( - self.llm_policy_cfg.pretrain_llm_path, - trust_remote_code=True, - torch_dtype=torch.bfloat16, # Use bfloat16 to save memory - device_map=None, # We'll manually move to device - ) - - # Apply LoRA if enabled - if self.llm_policy_cfg.use_lora: - logging.info("Applying LoRA for parameter-efficient fine-tuning") - lora_config = LoraConfig( - task_type=TaskType.CAUSAL_LM, - r=self.llm_policy_cfg.lora_r, - lora_alpha=self.llm_policy_cfg.lora_alpha, - lora_dropout=self.llm_policy_cfg.lora_dropout, - target_modules=["q_proj", "v_proj", "k_proj", "o_proj"], # Qwen-specific - ) - self.llm_policy_model = get_peft_model(self.llm_policy_model, lora_config) - self.llm_policy_model.print_trainable_parameters() - - self.llm_policy_model.to(self._cfg.device) - self.llm_policy_model.train() - - # ====================================================================== - # 3. [PRIORZERO-NEW] Initialize LLM Optimizer - # ====================================================================== - self._optimizer_llm = torch.optim.AdamW( - self.llm_policy_model.parameters(), - lr=self.llm_policy_cfg.llm_learning_rate, - weight_decay=self.llm_policy_cfg.llm_weight_decay, - betas=(0.9, 0.999), - ) - - # Optional: learning rate scheduler - self._lr_scheduler_llm = torch.optim.lr_scheduler.CosineAnnealingLR( - self._optimizer_llm, - T_max=100000, # Will be set from config - eta_min=self.llm_policy_cfg.llm_learning_rate * 0.1 - ) - - logging.info(f"✓ LLM Policy Model ({self.llm_policy_cfg.pretrain_llm_path}) initialized") - logging.info(f" - LLM learning rate: {self.llm_policy_cfg.llm_learning_rate}") - logging.info(f" - LoRA enabled: {self.llm_policy_cfg.use_lora}") - - # ====================================================================== - # 4. [PRIORZERO-NEW] Load Action Mappings - # ====================================================================== - if hasattr(self._cfg, 'action_map') and self._cfg.action_map is not None: - self.action_map = self._cfg.action_map - self.action_inv_map = {v: k for k, v in self.action_map.items()} - logging.info(f"✓ Action mappings loaded ({len(self.action_map)} actions)") - else: - logging.warning("⚠ Action mappings not found in config. Will use index-based actions.") - # Fallback: create dummy mappings - action_space_size = self._cfg.model.action_space_size - self.action_inv_map = {i: f"action_{i}" for i in range(action_space_size)} - self.action_map = {v: k for k, v in self.action_inv_map.items()} - - def _forward_learn(self, data: Tuple[torch.Tensor]) -> Dict[str, Union[float, int]]: - """ - [PRIORZERO-MODIFIED] - Dual-model training: UniZero world model + LLM policy. - - Training process: - 1. Train UniZero world model with standard losses (value, policy, reward, latent) - 2. Train LLM with SFT (supervised by MCTS policies) - 3. Optionally train LLM with RFT (reinforced by environment rewards) - 4. Joint optimization with combined loss - - Args: - data: Tuple containing (current_batch, target_batch, train_iter, game_segments) - - Returns: - log_dict: Dictionary of training metrics - """ - import logging - - self._learn_model.train() - self.llm_policy_model.train() - - # Unpack data - # NOTE: game_segments is our custom GameSegment with mcts_policy_segment - # [FIX] Handle both 3-element (from buffer) and 4-element (with explicit train_iter) formats - if len(data) == 4: - # Format: [current_batch, target_batch, train_iter, game_segments] - # This is when learner explicitly adds train_iter - current_batch, target_batch, train_iter, game_segments = data - elif len(data) == 3: - # Format: [current_batch, target_batch, game_segments] - # This is the standard format from PriorZeroGameBuffer.sample() - current_batch, target_batch, game_segments = data - train_iter = self._train_iteration # Get from instance variable - import logging - logger = logging.getLogger(__name__) - logger.debug( - f"[PRIORZERO] Using 3-element format. game_segments: " - f"{type(game_segments)}, count: {len(game_segments) if game_segments else 0}" - ) - else: - raise ValueError(f"Unexpected data format: expected 3 or 4 elements, got {len(data)}") - - # ============================================================================== - # Part 1: UniZero World Model Training (Full Implementation) - # ============================================================================== - - # Unpack batches - (obs_batch_ori, action_batch, mask_batch, batch_index_tensor, - weights, make_time) = current_batch[:6] - target_reward, target_value, target_policy = target_batch - - # Handle optional timestep - if len(current_batch) > 6: - timestep_batch = current_batch[6] - else: - timestep_batch = None - - # Convert to tensors and move to device - data_list = [mask_batch, target_reward, target_value, target_policy, weights] - (mask_batch, target_reward, target_value, - target_policy, weights) = to_torch_float_tensor(data_list, self._cfg.device) - - # Reshape targets - batch_size = self._cfg.batch_size - target_reward = target_reward.view(batch_size, -1) - target_value = target_value.view(batch_size, -1) - - # Apply scalar transform (for value and reward) - # [FIX] Use scalar_transform function (not self.scalar_transform) - # scalar_transform is a standalone function imported from lzero.policy - transformed_target_reward = scalar_transform(target_reward) - transformed_target_value = scalar_transform(target_value) - - # Convert to categorical distribution (for distributional RL) - target_reward_categorical = phi_transform( - self.reward_support, transformed_target_reward - ) - target_value_categorical = phi_transform( - self.value_support, transformed_target_value - ) - - # Prepare batch for world model - # NOTE: This follows the exact format required by UniZero world model - # [FIX] Convert obs_batch_ori to tensor if needed - if not isinstance(obs_batch_ori, torch.Tensor): - # [DEBUG] Check obs_batch_ori shape - import logging - logger = logging.getLogger(__name__) - if isinstance(obs_batch_ori, np.ndarray): - logger.info(f"[DEBUG] obs_batch_ori type: numpy, shape: {obs_batch_ori.shape}, dtype: {obs_batch_ori.dtype}") - - # [FIX] Reshape if observations are flattened (2D instead of 3D) - # Expected: [batch_size, num_unroll_steps+1, obs_dim] (buffer includes next_obs) - # Got: [batch_size, (num_unroll_steps+1) * obs_dim] - if len(obs_batch_ori.shape) == 2: - # Infer num_unroll_steps and obs_dim - # For text: obs_dim should be max_seq_len (e.g., 512) - obs_dim = 512 # Standard max_seq_len for BERT - total_size = obs_batch_ori.shape[1] - if total_size % obs_dim == 0: - inferred_steps = total_size // obs_dim - # Simply reshape to [batch_size, inferred_steps, obs_dim] - # The truncation to match action_batch will happen later (like unizero.py line 675) - obs_batch_ori = obs_batch_ori.reshape(batch_size, inferred_steps, obs_dim) - logger.info(f"[RESHAPE] Reshaped obs_batch_ori from (batch_size, {total_size}) to {obs_batch_ori.shape}") - else: - logger.warning(f"[RESHAPE_ERROR] Cannot reshape: total_size ({total_size}) not divisible by obs_dim ({obs_dim})") - - # Check if it's an object array (inhomogeneous shapes) - if obs_batch_ori.dtype == np.object_: - logger.warning(f"[SHAPE_ISSUE] obs_batch_ori is object array - inhomogeneous shapes!") - logger.warning(f"[SHAPE_ISSUE] First element shape: {obs_batch_ori[0].shape if len(obs_batch_ori) > 0 else 'N/A'}") - if len(obs_batch_ori) > 1: - logger.warning(f"[SHAPE_ISSUE] Second element shape: {obs_batch_ori[1].shape}") - # Try to handle inhomogeneous array by padding/truncating - # For now, just raise a descriptive error - raise ValueError( - f"obs_batch_ori has inhomogeneous shapes. " - f"First element shape: {obs_batch_ori[0].shape}, " - f"Cannot directly convert to tensor. " - f"This suggests the replay buffer is storing observations with different sequence lengths." - ) - obs_batch_ori = torch.from_numpy(obs_batch_ori).to(self._cfg.device) - - # [FIX] Convert action_batch to tensor and handle shape correctly - if not isinstance(action_batch, torch.Tensor): - action_batch = torch.from_numpy(action_batch).to(self._cfg.device) - - if action_batch.shape[-1] == 1: - actions_processed = action_batch.squeeze(-1).long() - elif len(action_batch.shape) == 1: - actions_processed = action_batch.long() - else: - actions_processed = action_batch.long() - - if timestep_batch is not None: - # Convert timestep_batch to tensor if needed - if not isinstance(timestep_batch, torch.Tensor): - timestep_batch = torch.from_numpy(timestep_batch).to(self._cfg.device) - - # Handle timestep_batch shape - if timestep_batch.shape[-1] == 1: - timestep_processed = timestep_batch.squeeze(-1).long() - elif len(timestep_batch.shape) == 1: - timestep_processed = timestep_batch.long() - else: - timestep_processed = timestep_batch.long() - - batch_for_gpt = { - 'observations': obs_batch_ori, - 'actions': actions_processed, - 'timestep': timestep_processed, - 'rewards': target_reward_categorical[:, :-1], - 'target_value': target_value_categorical[:, :-1], - 'target_policy': target_policy[:, :-1], - } - else: - batch_for_gpt = { - 'observations': obs_batch_ori, - 'actions': actions_processed, - 'rewards': target_reward_categorical[:, :-1], - 'target_value': target_value_categorical[:, :-1], - 'target_policy': target_policy[:, :-1], - } - - # [FIX] Following unizero.py lines 673-675 exactly: - # Convert mask_batch to boolean, then truncate to align with observations/rewards - batch_for_gpt['mask_padding'] = mask_batch == 1.0 # 0 means invalid padding data. Shape: (B, T) - - # [DEBUG] Log shapes before truncation - logger.info(f"[SHAPE_DEBUG] Before truncation: obs={batch_for_gpt['observations'].shape}, " - f"mask_padding={batch_for_gpt['mask_padding'].shape}, " - f"actions={batch_for_gpt['actions'].shape}") - - # [CRITICAL] Truncate observations to align with rewards/actions - # - observations from buffer include next_obs → shape (B, T+1, obs_dim) - # - mask_padding is already (B, T) from buffer - DO NOT truncate again! - # - After target processing: rewards[:, :-1] → (B, T-1) - # - So only observations need truncation - batch_for_gpt['observations'] = batch_for_gpt['observations'][:, :-1] # Shape: (B, T-1, obs_dim) - - # [FIX] Check if mask_padding needs truncation based on actual shape - if batch_for_gpt['mask_padding'].shape[1] > batch_for_gpt['observations'].shape[1]: - logger.warning(f"[SHAPE_FIX] Truncating mask_padding from {batch_for_gpt['mask_padding'].shape} to match obs") - batch_for_gpt['mask_padding'] = batch_for_gpt['mask_padding'][:, :-1] - - logger.info(f"[SHAPE_DEBUG] After truncation: obs={batch_for_gpt['observations'].shape}, " - f"mask_padding={batch_for_gpt['mask_padding'].shape}") - - # [FIX] Add missing 'ends' field (following unizero.py line 676) - # 'ends' marks terminal states in the trajectory (0 = not terminal) - batch_for_gpt['ends'] = torch.zeros(batch_for_gpt['mask_padding'].shape, dtype=torch.long, device=self._cfg.device) - - # [FIX] Add 'scalar_target_value' field for priority calculation (following unizero.py line 681) - batch_for_gpt['scalar_target_value'] = target_value - - # [FIX] Log shapes for debugging - import logging - logger = logging.getLogger(__name__) - logger.info(f"[BATCH_SHAPES] obs: {batch_for_gpt['observations'].shape}, actions: {batch_for_gpt['actions'].shape}, rewards: {batch_for_gpt['rewards'].shape}, mask_padding: {batch_for_gpt['mask_padding'].shape}") - - # Compute world model loss - wm_losses = self._learn_model.world_model.compute_loss( - batch_for_gpt, - self._target_model.world_model.tokenizer, - self.value_inverse_scalar_transform_handle, - ) - - # Weighted world model loss (for prioritized experience replay) - wm_total_loss = (weights * wm_losses.loss_total).mean() - - # ============================================================================== - # Part 2: [PRIORZERO-NEW] LLM Policy Training (SFT + RFT) - # ============================================================================== - - llm_sft_loss = torch.tensor(0.0, device=self._cfg.device) - llm_rft_loss = torch.tensor(0.0, device=self._cfg.device) - num_sft_samples = 0 - num_rft_samples = 0 - - # [FIX] Only perform LLM training if game_segments available - # [DEBUG] Always log game_segments status - logger = logging.getLogger(__name__) - logger.info(f"[LLM Training] game_segments type: {type(game_segments)}, " - f"is None: {game_segments is None}, " - f"len: {len(game_segments) if game_segments is not None else 'N/A'}") - - # [DEBUG] Check first segment's data - if game_segments is not None and len(game_segments) > 0: - seg0 = game_segments[0] - logger.info(f"[LLM Training] First segment stats: " - f"mcts_policies={len(seg0.mcts_policy_segment) if hasattr(seg0, 'mcts_policy_segment') else 0}, " - f"raw_obs={len([x for x in (seg0.raw_obs_segment if hasattr(seg0, 'raw_obs_segment') else []) if x is not None])}/{len(seg0.raw_obs_segment) if hasattr(seg0, 'raw_obs_segment') else 0}, " - f"actions={len(seg0.action_segment) if hasattr(seg0, 'action_segment') else 0}") - - if game_segments is not None and len(game_segments) > 0: - # Collect training data from game segments - sft_prompts = [] - sft_targets = [] - rft_prompts = [] - rft_rewards = [] - - # [DEBUG] Log segment information - logger.info(f"[LLM Training] Processing {len(game_segments)} game segments") - - for seg_idx, segment in enumerate(game_segments): - # [FIX] Use action_segment length, not obs_segment - # obs_segment includes frame_stack + unroll_steps, while - # mcts_policy_segment only has entries for actual actions taken - segment_length = len(segment.action_segment) - - # [FIX] Ensure mcts_policy_segment has the same length - # It might be a list or numpy array depending on whether game_segment_to_array() was called - mcts_policy_length = len(segment.mcts_policy_segment) if hasattr(segment, 'mcts_policy_segment') else 0 - - # [DEBUG] Log segment lengths for debugging - if self._cfg.get('debug_segment_processing', False): - obs_len = len(segment.obs_segment) if hasattr(segment, 'obs_segment') else 0 - raw_obs_len = len(segment.raw_obs_segment) if hasattr(segment, 'raw_obs_segment') else 0 - logging.info( - f"[Segment {seg_idx}] action_len={segment_length}, " - f"mcts_policy_len={mcts_policy_length}, obs_len={obs_len}, raw_obs_len={raw_obs_len}" - ) - - # [SAFETY] Use the minimum of the two lengths to avoid IndexError - max_index = min(segment_length, mcts_policy_length) - - if max_index == 0: - if self._cfg.get('debug_segment_processing', False): - logging.warning(f"[Segment {seg_idx}] Empty segment, skipping") - continue # Skip empty segments - - for i in range(max_index): - # [FIX] Safe access to mcts_policy_segment with bounds check - try: - mcts_policy = segment.mcts_policy_segment[i] - except (IndexError, KeyError, TypeError) as e: - # Log detailed error information for debugging - if self._cfg.get('debug_segment_processing', False): - logging.error( - f"[Segment {seg_idx}, Index {i}] Failed to access mcts_policy_segment: {e}\n" - f" segment_length={segment_length}, mcts_policy_length={mcts_policy_length}\n" - f" mcts_policy_segment type: {type(segment.mcts_policy_segment)}" - ) - continue - - # Skip if no MCTS policy available - if mcts_policy is None: - continue - - # [FIX] Use raw_obs_segment for text observations - # PriorZero's GameSegment stores raw text in raw_obs_segment - raw_obs_text = None - if hasattr(segment, 'raw_obs_segment') and i < len(segment.raw_obs_segment): - raw_obs_text = segment.raw_obs_segment[i] - elif i < len(segment.obs_segment): - # Fallback to obs_segment if raw_obs_segment not available - raw_obs_text = str(segment.obs_segment[i]) - - # Skip if raw_obs_text is None - if raw_obs_text is None: - continue - - # Build history context - history = [] - for j in range(max(0, i - self.llm_policy_cfg.history_length), i): - # [FIX] Use raw_obs_segment for history as well - obs_text = None - if hasattr(segment, 'raw_obs_segment') and j < len(segment.raw_obs_segment): - obs_text = segment.raw_obs_segment[j] - elif j < len(segment.obs_segment): - obs_text = str(segment.obs_segment[j]) - - if obs_text is not None and j < len(segment.action_segment): - history.append(( - obs_text, - self.action_inv_map.get(segment.action_segment[j], f"action_{segment.action_segment[j]}"), - float(segment.reward_segment[j]) if j < len(segment.reward_segment) else 0.0 - )) - - # Build prompt - instruction = build_llm_prompt( - current_obs=raw_obs_text, - history=history, - use_cot=self.llm_policy_cfg.use_cot - ) - - # Apply chat template - prompt = self.llm_tokenizer.apply_chat_template( - [{"role": "user", "content": instruction}], - tokenize=False, - add_generation_prompt=True - ) - - # ============================================================ - # SFT: Supervised Fine-Tuning with MCTS Policy - # ============================================================ - if self.llm_policy_cfg.sft_target == 'mcts_policy': - # [FIX] Use the mcts_policy we already safely retrieved above - # Don't access segment.mcts_policy_segment[i] again to avoid IndexError - mcts_policy_vec = mcts_policy - - # Convert MCTS policy to ranked action text - target_text = format_mcts_policy_to_text( - mcts_policy_vec, - self.action_inv_map, - top_k=5 - ) - - sft_prompts.append(prompt) - sft_targets.append(target_text) - num_sft_samples += 1 - - # ============================================================ - # RFT: Reinforcement Fine-Tuning with Environment Reward - # ============================================================ - if self.llm_policy_cfg.enable_rft and i < len(segment.reward_segment): - env_reward = float(segment.reward_segment[i]) - - # TODO - # Only use transitions with non-zero reward for RFT - if abs(env_reward) > 1e-9: - rft_prompts.append(prompt) - rft_rewards.append(env_reward) - num_rft_samples += 1 - - # ============================================================ - # Train LLM with SFT (with gradient accumulation for memory efficiency) - # ============================================================ - # num_sft_samples=0 # TODO - if num_sft_samples > 0: - # [PRIORZERO-OOM-FIX] Use micro-batching with gradient accumulation - micro_batch_size = self.llm_policy_cfg.llm_micro_batch_size - num_micro_batches = (num_sft_samples + micro_batch_size - 1) // micro_batch_size - accumulation_steps = self.llm_policy_cfg.llm_gradient_accumulation_steps - - # Prepare full texts (prompt + target + eos) - full_texts = [ - p + t + self.llm_tokenizer.eos_token - for p, t in zip(sft_prompts, sft_targets) - ] - - # Process in micro-batches - accumulated_sft_loss = 0.0 - for micro_batch_idx in range(num_micro_batches): - start_idx = micro_batch_idx * micro_batch_size - end_idx = min((micro_batch_idx + 1) * micro_batch_size, num_sft_samples) - - # Get micro-batch - micro_batch_texts = full_texts[start_idx:end_idx] - micro_batch_prompts = sft_prompts[start_idx:end_idx] - - # Tokenize micro-batch - inputs = self.llm_tokenizer( - micro_batch_texts, - padding=True, - truncation=True, - max_length=self.llm_policy_cfg.prompt_max_len, - return_tensors="pt" - ).to(self._cfg.device) - - # Create labels (mask prompt tokens to only compute loss on target) - labels = inputs.input_ids.clone() - labels[labels == self.llm_tokenizer.pad_token_id] = -100 - - # Mask prompt tokens - for i, prompt in enumerate(micro_batch_prompts): - prompt_tokens = self.llm_tokenizer.encode(prompt, add_special_tokens=False) - prompt_len = len(prompt_tokens) - labels[i, :prompt_len] = -100 - - # Forward pass - llm_outputs = self.llm_policy_model( - input_ids=inputs.input_ids, - attention_mask=inputs.attention_mask, - labels=labels - ) - - # Scale loss by number of accumulation steps (for correct gradient magnitude) - micro_batch_loss = llm_outputs.loss / accumulation_steps - accumulated_sft_loss += micro_batch_loss.item() - - # Backward pass (accumulate gradients) - micro_batch_loss.backward() - - # Free memory - del inputs, labels, llm_outputs - torch.cuda.empty_cache() - - # Average loss for logging - llm_sft_loss = torch.tensor(accumulated_sft_loss, device=self._cfg.device) - - # ============================================================ - # Train LLM with RFT (Policy Gradient with gradient accumulation) - # ============================================================ - if num_rft_samples > 0 and self.llm_policy_cfg.enable_rft: - # [PRIORZERO-OOM-FIX] Use micro-batching with gradient accumulation - micro_batch_size = self.llm_policy_cfg.llm_micro_batch_size - num_micro_batches = (num_rft_samples + micro_batch_size - 1) // micro_batch_size - accumulation_steps = self.llm_policy_cfg.llm_gradient_accumulation_steps - - # Process in micro-batches - accumulated_rft_loss = 0.0 - for micro_batch_idx in range(num_micro_batches): - start_idx = micro_batch_idx * micro_batch_size - end_idx = min((micro_batch_idx + 1) * micro_batch_size, num_rft_samples) - - # Get micro-batch - micro_batch_prompts = rft_prompts[start_idx:end_idx] - micro_batch_rewards = rft_rewards[start_idx:end_idx] - - # Tokenize prompts - inputs = self.llm_tokenizer( - micro_batch_prompts, - padding=True, - truncation=True, - max_length=self.llm_policy_cfg.prompt_max_len, - return_tensors="pt" - ).to(self._cfg.device) - - # [FIX] Forward pass WITH gradient tracking (remove no_grad) - outputs = self.llm_policy_model( - input_ids=inputs.input_ids, - attention_mask=inputs.attention_mask - ) - - # Compute policy gradient loss (REINFORCE) - # Loss = -reward * log_prob(action) - logits = outputs.logits - log_probs = F.log_softmax(logits, dim=-1) - - # Get log probability of actual tokens - shifted_log_probs = log_probs[:, :-1, :].contiguous() - shifted_labels = inputs.input_ids[:, 1:].contiguous() - - # Gather log probs of actual tokens - token_log_probs = shifted_log_probs.gather( - dim=-1, - index=shifted_labels.unsqueeze(-1) - ).squeeze(-1) - - # Mask padding tokens - mask = (shifted_labels != self.llm_tokenizer.pad_token_id).float() - token_log_probs = token_log_probs * mask - - # Sum log probs per sequence - sequence_log_probs = token_log_probs.sum(dim=-1) / (mask.sum(dim=-1) + 1e-8) - - # Compute REINFORCE loss for micro-batch - rewards_tensor = torch.tensor( - micro_batch_rewards, - device=self._cfg.device, - dtype=torch.float32 - ) - - # Normalize rewards within micro-batch (important for stable training) - if len(micro_batch_rewards) > 1: - rewards_tensor = (rewards_tensor - rewards_tensor.mean()) / (rewards_tensor.std() + 1e-8) - - micro_batch_rft_loss = -(rewards_tensor * sequence_log_probs).mean() / accumulation_steps - accumulated_rft_loss += micro_batch_rft_loss.item() - - # Backward pass (accumulate gradients) - micro_batch_rft_loss.backward() - - # Free memory - del inputs, outputs, logits, log_probs, rewards_tensor - torch.cuda.empty_cache() - - # Average loss for logging - llm_rft_loss = torch.tensor(accumulated_rft_loss, device=self._cfg.device) - - # ============================================================================== - # Part 3: Joint Optimization - # ============================================================================== - - # [PRIORZERO-OOM-FIX] Note: LLM gradients already accumulated via micro-batching above - # Only need to compute world model gradients here - - # Combine losses (for logging only - LLM loss already backpropagated) - llm_loss = ( - self.llm_policy_cfg.llm_loss_weight * llm_sft_loss + - self.llm_policy_cfg.rft_loss_weight * llm_rft_loss - ) - total_loss = wm_total_loss + llm_loss # For logging - - # Zero world model gradients only (LLM gradients already accumulated) - self._optimizer_world_model.zero_grad() - - # Backward pass for world model only - wm_total_loss.backward() - - # Gradient clipping for both models - wm_grad_norm = torch.nn.utils.clip_grad_norm_( - self._learn_model.world_model.parameters(), - self._cfg.grad_clip_value - ) - llm_grad_norm = torch.nn.utils.clip_grad_norm_( - self.llm_policy_model.parameters(), - self._cfg.grad_clip_value - ) - - # Optimizer step for both models - self._optimizer_world_model.step() - self._optimizer_llm.step() # Apply accumulated LLM gradients - - # Zero LLM gradients after step (ready for next iteration) - self._optimizer_llm.zero_grad() - - # Learning rate scheduler step (optional) - if self._lr_scheduler_llm is not None: - self._lr_scheduler_llm.step() - - # Update target model (soft update) - self._target_model.update(self._learn_model.state_dict()) - - # ============================================================================== - # Part 4: Logging (Aligned with UniZero) - # ============================================================================== - - # Extract intermediate losses from world model (like UniZero) - intermediate_losses = wm_losses.intermediate_losses - obs_loss = intermediate_losses.get('loss_obs', torch.tensor(0.0)) - reward_loss = intermediate_losses.get('loss_rewards', torch.tensor(0.0)) - policy_loss = intermediate_losses.get('loss_policy', torch.tensor(0.0)) - value_loss = intermediate_losses.get('loss_value', torch.tensor(0.0)) - latent_recon_loss = intermediate_losses.get('latent_recon_loss', torch.tensor(0.0)) - perceptual_loss = intermediate_losses.get('perceptual_loss', torch.tensor(0.0)) - orig_policy_loss = intermediate_losses.get('orig_policy_loss', torch.tensor(0.0)) - policy_entropy = intermediate_losses.get('policy_entropy', torch.tensor(0.0)) - first_step_losses = intermediate_losses.get('first_step_losses', {}) - middle_step_losses = intermediate_losses.get('middle_step_losses', {}) - last_step_losses = intermediate_losses.get('last_step_losses', {}) - - # Analysis metrics (dormant ratio, weight magnitude, etc.) - dormant_ratio_encoder = intermediate_losses.get('dormant_ratio_encoder', 0.0) - dormant_ratio_transformer = intermediate_losses.get('dormant_ratio_transformer', 0.0) - dormant_ratio_head = intermediate_losses.get('dormant_ratio_head', 0.0) - avg_weight_mag_encoder = intermediate_losses.get('avg_weight_mag_encoder', 0.0) - avg_weight_mag_transformer = intermediate_losses.get('avg_weight_mag_transformer', 0.0) - avg_weight_mag_head = intermediate_losses.get('avg_weight_mag_head', 0.0) - e_rank_last_linear = intermediate_losses.get('e_rank_last_linear', 0.0) - e_rank_sim_norm = intermediate_losses.get('e_rank_sim_norm', 0.0) - latent_state_l2_norms = intermediate_losses.get('latent_state_l2_norms', torch.tensor(0.0)) - latent_action_l2_norms = intermediate_losses.get('latent_action_l2_norms', 0.0) - - # Logits statistics - logits_value_mean = intermediate_losses.get('logits_value_mean', 0.0) - logits_value_max = intermediate_losses.get('logits_value_max', 0.0) - logits_value_min = intermediate_losses.get('logits_value_min', 0.0) - logits_policy_mean = intermediate_losses.get('logits_policy_mean', 0.0) - logits_policy_max = intermediate_losses.get('logits_policy_max', 0.0) - logits_policy_min = intermediate_losses.get('logits_policy_min', 0.0) - - # Temperature parameters - temperature_value = intermediate_losses.get('temperature_value', 0.0) - temperature_reward = intermediate_losses.get('temperature_reward', 0.0) - temperature_policy = intermediate_losses.get('temperature_policy', 0.0) - - # Value priority for prioritized replay - value_priority_tensor = intermediate_losses.get('value_priority', torch.tensor([0.0])) - value_priority_np = value_priority_tensor.detach().cpu().numpy() + 1e-6 - - # Compute target policy entropy (for analysis) - valid_target_policy = batch_for_gpt['target_policy'][batch_for_gpt['mask_padding']] - target_policy_entropy = -torch.sum(valid_target_policy * torch.log(valid_target_policy + 1e-9), dim=-1) - average_target_policy_entropy = target_policy_entropy.mean() - - # Build comprehensive log dict (aligned with UniZero) - log_dict = { - # ============ Core Losses ============ - 'weighted_total_loss': wm_total_loss.item(), - 'obs_loss': obs_loss.item() if torch.is_tensor(obs_loss) else obs_loss, - 'reward_loss': reward_loss.item() if torch.is_tensor(reward_loss) else reward_loss, - 'policy_loss': policy_loss.item() if torch.is_tensor(policy_loss) else policy_loss, - 'value_loss': value_loss.item() if torch.is_tensor(value_loss) else value_loss, - 'latent_recon_loss': latent_recon_loss.item() if torch.is_tensor(latent_recon_loss) else latent_recon_loss, - 'perceptual_loss': perceptual_loss.item() if torch.is_tensor(perceptual_loss) else perceptual_loss, - 'orig_policy_loss': orig_policy_loss.item() if torch.is_tensor(orig_policy_loss) else orig_policy_loss, - 'policy_entropy': policy_entropy.item() if torch.is_tensor(policy_entropy) else policy_entropy, - 'target_policy_entropy': average_target_policy_entropy.item(), - - - # ============ Step-wise Losses ============ - 'analysis/first_step_loss_value': first_step_losses.get('loss_value', torch.tensor(0.0)).item() if isinstance(first_step_losses.get('loss_value'), torch.Tensor) else 0.0, - 'analysis/first_step_loss_policy': first_step_losses.get('loss_policy', torch.tensor(0.0)).item() if isinstance(first_step_losses.get('loss_policy'), torch.Tensor) else 0.0, - 'analysis/first_step_loss_rewards': first_step_losses.get('loss_rewards', torch.tensor(0.0)).item() if isinstance(first_step_losses.get('loss_rewards'), torch.Tensor) else 0.0, - 'analysis/first_step_loss_obs': first_step_losses.get('loss_obs', torch.tensor(0.0)).item() if isinstance(first_step_losses.get('loss_obs'), torch.Tensor) else 0.0, - - 'analysis/middle_step_loss_value': middle_step_losses.get('loss_value', torch.tensor(0.0)).item() if isinstance(middle_step_losses.get('loss_value'), torch.Tensor) else 0.0, - 'analysis/middle_step_loss_policy': middle_step_losses.get('loss_policy', torch.tensor(0.0)).item() if isinstance(middle_step_losses.get('loss_policy'), torch.Tensor) else 0.0, - 'analysis/middle_step_loss_rewards': middle_step_losses.get('loss_rewards', torch.tensor(0.0)).item() if isinstance(middle_step_losses.get('loss_rewards'), torch.Tensor) else 0.0, - 'analysis/middle_step_loss_obs': middle_step_losses.get('loss_obs', torch.tensor(0.0)).item() if isinstance(middle_step_losses.get('loss_obs'), torch.Tensor) else 0.0, - - 'analysis/last_step_loss_value': last_step_losses.get('loss_value', torch.tensor(0.0)).item() if isinstance(last_step_losses.get('loss_value'), torch.Tensor) else 0.0, - 'analysis/last_step_loss_policy': last_step_losses.get('loss_policy', torch.tensor(0.0)).item() if isinstance(last_step_losses.get('loss_policy'), torch.Tensor) else 0.0, - 'analysis/last_step_loss_rewards': last_step_losses.get('loss_rewards', torch.tensor(0.0)).item() if isinstance(last_step_losses.get('loss_rewards'), torch.Tensor) else 0.0, - 'analysis/last_step_loss_obs': last_step_losses.get('loss_obs', torch.tensor(0.0)).item() if isinstance(last_step_losses.get('loss_obs'), torch.Tensor) else 0.0, - - # ============ Analysis Metrics ============ - 'analysis/dormant_ratio_encoder': dormant_ratio_encoder, - 'analysis/dormant_ratio_transformer': dormant_ratio_transformer, - 'analysis/dormant_ratio_head': dormant_ratio_head, - 'analysis/avg_weight_mag_encoder': avg_weight_mag_encoder, - 'analysis/avg_weight_mag_transformer': avg_weight_mag_transformer, - 'analysis/avg_weight_mag_head': avg_weight_mag_head, - 'analysis/e_rank_last_linear': e_rank_last_linear, - 'analysis/e_rank_sim_norm': e_rank_sim_norm, - 'analysis/latent_state_l2_norms': latent_state_l2_norms.item() if torch.is_tensor(latent_state_l2_norms) else latent_state_l2_norms, - 'analysis/latent_action_l2_norms': latent_action_l2_norms, - - # ============ Logits Statistics ============ - 'logits_value_mean': logits_value_mean, - 'logits_value_max': logits_value_max, - 'logits_value_min': logits_value_min, - 'logits_policy_mean': logits_policy_mean, - 'logits_policy_max': logits_policy_max, - 'logits_policy_min': logits_policy_min, - - # ============ Temperature Parameters ============ - 'temperature_value': temperature_value, - 'temperature_reward': temperature_reward, - 'temperature_policy': temperature_policy, - - # ============ Targets ============ - 'target_reward': target_reward.mean().item(), - 'target_value': target_value.mean().item(), - 'transformed_target_reward': transformed_target_reward.mean().item(), - 'transformed_target_value': transformed_target_value.mean().item(), - 'value_priority': value_priority_np.mean().item(), - 'value_priority_orig': value_priority_np, - - # ============ Gradient Norms ============ - 'total_grad_norm_before_clip_wm': wm_grad_norm.item(), - 'llm_grad_norm': llm_grad_norm.item(), - - # ============ Learning Rates ============ - 'cur_lr_world_model': self._optimizer_world_model.param_groups[0]['lr'], - 'llm_lr': self._optimizer_llm.param_groups[0]['lr'], - - # ============ [PRIORZERO] LLM-specific Metrics ============ - 'llm_sft_loss': llm_sft_loss.item(), - 'llm_rft_loss': llm_rft_loss.item(), - 'llm_total_loss': llm_loss.item(), - 'num_sft_samples': float(num_sft_samples), - 'num_rft_samples': float(num_rft_samples), - 'total_loss': total_loss.item(), - } - - # ============================================================================== - # [PRIORZERO-NEW] WandB Logging (if enabled) - # ============================================================================== - if self._cfg.get('use_wandb', False): - try: - import wandb - if wandb.run is not None: - # Log all metrics to WandB with hierarchical naming - wandb.log({ - # World Model Metrics - 'train/wm/total_loss': log_dict['wm_total_loss'], - 'train/wm/value_loss': log_dict['wm_value_loss'], - 'train/wm/policy_loss': log_dict['wm_policy_loss'], - 'train/wm/reward_loss': log_dict['wm_reward_loss'], - 'train/wm/grad_norm': log_dict['wm_grad_norm'], - 'train/wm/learning_rate': log_dict['wm_lr'], - - # LLM Policy Metrics - 'train/llm/sft_loss': log_dict['llm_sft_loss'], - 'train/llm/rft_loss': log_dict['llm_rft_loss'], - 'train/llm/total_loss': log_dict['llm_total_loss'], - 'train/llm/grad_norm': log_dict['llm_grad_norm'], - 'train/llm/learning_rate': log_dict['llm_lr'], - 'train/llm/num_sft_samples': float(log_dict['num_sft_samples']), - 'train/llm/num_rft_samples': float(log_dict['num_rft_samples']), - - # Combined Metrics - 'train/total_loss': log_dict['total_loss'], - }, step=self._train_iteration) - except Exception as e: - # Don't fail training if wandb logging fails - import logging - logging.warning(f"WandB logging failed: {e}") - - return log_dict - - def _monitor_vars_learn(self) -> List[str]: - """ - [PRIORZERO-MODIFIED] - Register variables to be monitored in learn mode for TensorBoard logging. - - This extends UniZero's monitoring with PriorZero-specific LLM metrics. - - Returns: - List of variable names that should be logged to TensorBoard/WandB - """ - - return [ - # ============ LLM Loss Metrics ============ - 'llm_sft_loss', # Supervised fine-tuning loss - 'llm_rft_loss', # Reinforcement fine-tuning loss - 'llm_total_loss', # Combined LLM loss - 'llm_grad_norm', # LLM gradient norm - 'llm_lr', # LLM learning rate - - # ============ LLM Training Statistics ============ - 'num_sft_samples', # Number of SFT samples in batch - 'num_rft_samples', # Number of RFT samples in batch - - # ============ Combined Metrics ============ - 'total_loss', # Total loss (WM + LLM) - 'wm_total_loss', # World model total loss - 'wm_grad_norm', # World model gradient norm - 'wm_lr', # World model learning rate - - # ============ World Model Component Losses ============ - 'wm_value_loss', - 'wm_policy_loss', - 'wm_reward_loss', - 'wm_obs_loss', - - 'analysis/dormant_ratio_encoder', - 'analysis/dormant_ratio_transformer', - 'analysis/dormant_ratio_head', - - 'analysis/avg_weight_mag_encoder', - 'analysis/avg_weight_mag_transformer', - 'analysis/avg_weight_mag_head', - 'analysis/e_rank_last_linear', - 'analysis/e_rank_sim_norm', - - 'analysis/latent_state_l2_norms', - 'analysis/l2_norm_before', - 'analysis/l2_norm_after', - 'analysis/grad_norm_before', - 'analysis/grad_norm_after', - - 'analysis/first_step_loss_value', - 'analysis/first_step_loss_policy', - 'analysis/first_step_loss_rewards', - 'analysis/first_step_loss_obs', - - 'analysis/middle_step_loss_value', - 'analysis/middle_step_loss_policy', - 'analysis/middle_step_loss_rewards', - 'analysis/middle_step_loss_obs', - - 'analysis/last_step_loss_value', - 'analysis/last_step_loss_policy', - 'analysis/last_step_loss_rewards', - 'analysis/last_step_loss_obs', - - 'adaptive_alpha', - "adaptive_target_entropy_ratio", - 'alpha_loss', - - 'Current_GPU', - 'Max_GPU', - 'collect_epsilon', - 'collect_mcts_temperature', - 'cur_lr_world_model', - 'cur_lr_tokenizer', - - 'weighted_total_loss', - 'obs_loss', - 'policy_loss', - 'orig_policy_loss', - 'policy_entropy', - 'latent_recon_loss', - 'target_policy_entropy', - 'reward_loss', - 'value_loss', - 'consistency_loss', - 'value_priority', - 'target_reward', - 'target_value', - 'total_grad_norm_before_clip_wm', - # tokenizer - 'commitment_loss', - 'reconstruction_loss', - 'perceptual_loss', - - - "logits_value_mean", - "logits_value_max", - "logits_value_min", - "logits_policy_mean", - "logits_policy_max", - "logits_policy_min", - - "temperature_value", - "temperature_reward", - "temperature_policy", - "current_policy_label_eps", - 'adaptive_alpha', - "adaptive_target_entropy_ratio", - 'alpha_loss', - "current_encoder_clip_value", - - # ==================== [新增] 添加范数和中间张量监控变量 ==================== - # 模块总范数 - 'norm/encoder/_total_norm', - 'norm/transformer/_total_norm', - 'norm/head_value/_total_norm', - 'norm/head_reward/_total_norm', - 'norm/head_policy/_total_norm', - # 中间张量 x 的统计信息 - 'norm/x_token/mean', - 'norm/x_token/std', - 'norm/x_token/max', - 'norm/x_token/min', - ] - # 注意:我们不把每一层的范数都加到这里,因为数量太多会导致日志混乱。 - # 在实践中,如果通过总范数发现问题,可以临时在TensorBoard中搜索特定层的范数, - # 或者在本地打印 `norm_log_dict` 来进行详细分析。 - # wandb等工具可以更好地处理大量的动态指标。 - # ======================================================================== - - - def _forward_collect( - self, - data: torch.Tensor, - action_mask: List[np.ndarray], - temperature: float = 1.0, - to_play: List[int] = None, - epsilon: float = 0.0, - ready_env_id: List[int] = None, - **kwargs - ) -> Dict[int, Dict[str, Any]]: - """ - [PRIORZERO-MODIFIED] - Forward pass for data collection with LLM-guided MCTS. - - Process: - 1. Get LLM prior outputs from kwargs - 2. Parse LLM outputs into policy priors - 3. Run world model initial inference - 4. Inject LLM priors into MCTS root node (replace policy logits) - 5. Run MCTS search with LLM-guided priors - 6. Return best action and statistics - - Args: - data: Stacked observations (tensor) - action_mask: Action masks for each environment - temperature: Temperature for action selection - to_play: Player IDs (for multi-agent) - epsilon: Epsilon for epsilon-greedy exploration - ready_env_id: List of ready environment IDs - **kwargs: Additional arguments, including 'llm_prior_outputs' - - Returns: - output_dict: Dictionary mapping env_id to action and search statistics - """ - self._collect_model.eval() - - # ====================================================================== - # [PRIORZERO-NEW] Get LLM Prior Outputs - # ====================================================================== - llm_prior_outputs = kwargs.pop('llm_prior_outputs', None) - - if llm_prior_outputs is None: - # If no LLM prior available, fall back to standard UniZero behavior - logging.debug("No LLM priors provided, using standard UniZero MCTS") - return super()._forward_collect( - data, action_mask, temperature, to_play, epsilon, - ready_env_id=ready_env_id, **kwargs - ) - - # ====================================================================== - # Parse LLM Outputs into Policy Priors - # ====================================================================== - policy_priors = [] - for output in llm_prior_outputs: - # Extract generated text - generated_text = output.outputs[0].text if hasattr(output, 'outputs') else str(output) - - # Parse into policy distribution - prior_policy = parse_llm_action_ranking( - generated_text, - self.action_map, - self._cfg.model.action_space_size, - fallback_to_uniform=True - ) - - # Convert to log probabilities (for compatibility with MCTS) - policy_logits = torch.log(torch.from_numpy(prior_policy) + 1e-9) - policy_priors.append(policy_logits) - - policy_priors = torch.stack(policy_priors).to(self._cfg.device) - - # ====================================================================== - # World Model Initial Inference - # ====================================================================== - with torch.no_grad(): - # Run representation network to get latent state - network_output = self._collect_model.initial_inference(data) - - # Unpack network outputs - latent_state_roots, reward_roots, pred_values, policy_logits_roots = \ - mz_network_output_unpack(network_output) - - # [PRIORZERO-KEY] Replace policy logits with LLM priors - network_output.policy_logits = policy_priors - - # Prepare for MCTS - if not self._cfg.mcts_ctree: - # Python implementation (not recommended for performance) - raise NotImplementedError("Python MCTS not supported for PriorZero") - - # ====================================================================== - # MCTS Search with LLM-Guided Priors - # ====================================================================== - # This is the key part where LLM priors guide the search - - # [FIX] Align with UniZero: construct legal_actions from action_mask - active_collect_env_num = len(ready_env_id) - legal_actions = [[i for i, x in enumerate(action_mask[j]) if x == 1] - for j in range(active_collect_env_num)] - - # Get timestep if available - timestep = kwargs.get('timestep', None) - - # [FIX] Align with UniZero: transform values and prepare data - pred_values_np = self.value_inverse_scalar_transform_handle(pred_values).detach().cpu().numpy() - latent_state_roots_np = latent_state_roots.detach().cpu().numpy() - # reward_roots_np = reward_roots.detach().cpu().numpy() - policy_logits_for_mcts = policy_priors.detach().cpu().numpy().tolist() - - # [FIX] Align with UniZero: Create MCTS roots with legal_actions (not action_space_size) - roots = MCTSCtree.roots(active_collect_env_num, legal_actions) - - # [FIX] Align with UniZero: noises based on number of valid actions per environment - noises = [ - np.random.dirichlet([self._cfg.root_dirichlet_alpha] * int(sum(action_mask[j])) - ).astype(np.float32).tolist() - for j in range(active_collect_env_num) - ] - - # [FIX] Align with UniZero: prepare roots (note reward_roots_np, not list(pred_values_np)) - roots.prepare( - self._cfg.root_noise_weight, - noises, - reward_roots, - # reward_roots_np, - policy_logits_for_mcts, - to_play if to_play is not None else [-1] * active_collect_env_num, - ) - - # Run MCTS search - MCTSCtree(self._cfg).search( - roots, - self._collect_model, - latent_state_roots_np, - reward_roots, - to_play if to_play is not None else [-1] * latent_state_roots_np.shape[0], - ) - - # Extract search results - roots_visit_count = roots.get_distributions() - roots_values = roots.get_values() - - # ====================================================================== - # [PRIORZERO] Get valid_actions_list for dynamic action mapping - # ====================================================================== - valid_actions_list = kwargs.get('valid_actions_list', None) - - # ====================================================================== - # Select Actions and Prepare Output (Aligned with UniZero) - # ====================================================================== - output = {} - - for i, env_id in enumerate(ready_env_id): - # [FIX] Get visit count distribution (only contains legal actions) - distributions = roots_visit_count[i] - value = roots_values[i] - - # [FIX] Use select_action from UniZero (aligns with UniZero line 1115-1117) - # NOTE: Only legal actions possess visit counts, so action_index_in_legal_action_set - # represents the index within the legal action set, not the entire action set - action_index_in_legal_action_set, visit_count_distribution_entropy = select_action( - distributions, - temperature=temperature if temperature is not None else self._collect_mcts_temperature, - deterministic=False - ) - - # [FIX] Convert action_index_in_legal_action_set to the actual action in full action space - # (aligns with UniZero line 1119) - legal_action_indices = np.where(action_mask[i] == 1.0)[0] - action = legal_action_indices[action_index_in_legal_action_set] - - # [PRIORZERO] Create dynamic action_inv_map for this specific state - # This maps action_index -> action_text using the current state's valid_actions - if valid_actions_list is not None and i < len(valid_actions_list): - dynamic_action_inv_map = { - idx: act_text - for idx, act_text in enumerate(valid_actions_list[i]) - } - else: - # Fallback to static mapping if valid_actions not available - dynamic_action_inv_map = self.action_inv_map - - output[env_id] = { - 'action': int(action), - 'visit_count_distributions': distributions, - 'visit_count_distribution_entropy': visit_count_distribution_entropy, - 'searched_value': value, - 'predicted_value': pred_values_np[i], - 'dynamic_action_inv_map': dynamic_action_inv_map, # [PRIORZERO] Include dynamic mapping - } - - return output - - def _state_dict_learn(self) -> Dict[str, Any]: - """ - [PRIORZERO-MODIFIED] - Save state dict for both world model and LLM. - """ - state_dict = super()._state_dict_learn() - - # Add LLM model and optimizer - state_dict['llm_model'] = self.llm_policy_model.state_dict() - state_dict['optimizer_llm'] = self._optimizer_llm.state_dict() - - if self._lr_scheduler_llm is not None: - state_dict['lr_scheduler_llm'] = self._lr_scheduler_llm.state_dict() - - return state_dict - - def _load_state_dict_learn(self, state_dict: Dict[str, Any]) -> None: - """ - [PRIORZERO-MODIFIED] - Load state dict for both world model and LLM. - """ - super()._load_state_dict_learn(state_dict) - - # Load LLM model and optimizer - if 'llm_model' in state_dict: - self.llm_policy_model.load_state_dict(state_dict['llm_model']) - logging.info("✓ LLM model state loaded") - - if 'optimizer_llm' in state_dict: - self._optimizer_llm.load_state_dict(state_dict['optimizer_llm']) - logging.info("✓ LLM optimizer state loaded") - - if 'lr_scheduler_llm' in state_dict and self._lr_scheduler_llm is not None: - self._lr_scheduler_llm.load_state_dict(state_dict['lr_scheduler_llm']) - logging.info("✓ LLM scheduler state loaded") diff --git a/zoo/jericho/priorzero/priorzero_prompts.py b/zoo/jericho/priorzero/priorzero_prompts.py deleted file mode 100644 index 4ce7ef787..000000000 --- a/zoo/jericho/priorzero/priorzero_prompts.py +++ /dev/null @@ -1,399 +0,0 @@ -""" -PriorZero LLM Prompts Module - -This module provides optimized prompt templates for PriorZero's LLM policy, -based on the successful prompt structure from Open-Reasoner-Zero. - -Key Features: -- Structured reasoning with and tags -- Clear role definitions (User/Assistant paradigm) -- Explicit format examples to guide the LLM -- Game-specific context integration - -Author: PriorZero Team -Date: 2025-10-21 -""" - -from jinja2 import Template -from typing import List, Dict, Any, Optional - - -class PriorZeroPromptTemplates: - """ - Centralized prompt templates for PriorZero LLM policy. - - Prompt Structure: - 1. System instruction (role definition) - 2. Format specification ( and tags) - 3. Example format to prime the model - 4. User query with game state - 5. Start reasoning with "" tag - """ - - # ============================================================================== - # MCTS Policy Guidance Prompts - # ============================================================================== - - MCTS_POLICY_TEMPLATE = """\ -{{bos_token}}A conversation between User and Assistant. The User is playing a text adventure game \ -and needs to decide the next action. The Assistant carefully analyzes the current game state, \ -considers the available actions, and recommends the best action to take. \ -The reasoning process is enclosed within tags, and the recommended action \ -is enclosed within tags. For example: \ - The player is in a dark room and needs light. The lamp is available. \ - take lamp . \ - -User: Current game state: -{{game_state}} - -Available actions: -{{valid_actions}} - -Recent history: -{{history}} - -What is the best action to take? -Assistant: \ -""" - - # ============================================================================== - # Supervised Fine-Tuning (SFT) Prompts - Learning from MCTS Policy - # ============================================================================== - - SFT_FROM_MCTS_TEMPLATE = """\ -{{bos_token}}A conversation between User and Assistant. The User is playing a text adventure game. \ -The Assistant provides step-by-step reasoning and selects the best action based on MCTS search results. \ -The reasoning is in tags and the action is in tags. \ - -User: Game state: {{game_state}} -Available actions: {{valid_actions}} -MCTS recommended action: {{mcts_action}} -MCTS value estimate: {{mcts_value}} - -Please explain why this is the best action and then select it. -Assistant: \ -""" - - # ============================================================================== - # Reward Fine-Tuning (RFT) Prompts - Learning from Environment Rewards - # ============================================================================== - - RFT_TEMPLATE = """\ -{{bos_token}}A conversation between User and Assistant. The User is playing a text adventure game \ -and wants to maximize the total reward. The Assistant analyzes the game state, considers past rewards, \ -and selects actions that lead to higher rewards. \ -The reasoning is in tags and the action is in tags. \ - -User: Current game state: -{{game_state}} - -Available actions: -{{valid_actions}} - -Recent trajectory: -{{trajectory_with_rewards}} - -Cumulative reward so far: {{cumulative_reward}} - -What action should I take to maximize future rewards? -Assistant: \ -""" - - # ============================================================================== - # Evaluation Prompts - For Testing LLM Policy - # ============================================================================== - - EVAL_TEMPLATE = """\ -{{bos_token}}A conversation between User and Assistant. The User is playing a text adventure game. \ -The Assistant thinks carefully about the situation and provides the best action. \ -Format: reasoning action . \ - -User: {{game_state}} -Available actions: {{valid_actions}} -Assistant: \ -""" - - # ============================================================================== - # Few-Shot Learning Prompts - With Example Demonstrations - # ============================================================================== - - FEW_SHOT_TEMPLATE = """\ -{{bos_token}}A conversation between User and Assistant. The User is playing a text adventure game. \ -The Assistant learns from examples and applies similar reasoning to new situations. \ - -Example 1: -User: You are in a dark room. You can't see anything. -Available actions: [go north, take lamp, light lamp] -Assistant: I need light to see. I should take the lamp first, then light it. take lamp - -Example 2: -User: You are holding a lamp. It is dark. -Available actions: [go north, light lamp, drop lamp] -Assistant: I have the lamp but it's not lit. I should light it to see. light lamp - -Now your turn: -User: {{game_state}} -Available actions: {{valid_actions}} -Assistant: \ -""" - - -class PriorZeroPromptBuilder: - """ - Builder class for constructing prompts with specific game context. - """ - - def __init__(self, tokenizer): - """ - Initialize the prompt builder. - - Args: - tokenizer: HuggingFace tokenizer with bos_token - """ - self.tokenizer = tokenizer - self.templates = PriorZeroPromptTemplates() - - def _get_bos_token(self) -> str: - """Get the beginning-of-sequence token.""" - if self.tokenizer.bos_token_id is None: - return "" - return self.tokenizer.decode([self.tokenizer.bos_token_id]) - - def build_mcts_policy_prompt( - self, - game_state: str, - valid_actions: List[str], - history: Optional[List[Dict[str, Any]]] = None, - ) -> str: - """ - Build a prompt for MCTS policy guidance. - - Args: - game_state: Current observation text from the game - valid_actions: List of valid action strings - history: Recent trajectory [(obs, action, reward), ...] - - Returns: - Formatted prompt string - """ - # Format valid actions as a numbered list - actions_str = "\n".join([f"{i+1}. {action}" for i, action in enumerate(valid_actions)]) - - # Format history - if history is None or len(history) == 0: - history_str = "This is the beginning of the game." - else: - history_lines = [] - for i, step in enumerate(history[-5:]): # Last 5 steps - obs = step.get('observation', 'N/A') - action = step.get('action', 'N/A') - reward = step.get('reward', 0) - history_lines.append(f"Step {i+1}: Observation: {obs[:100]}... | Action: {action} | Reward: {reward}") - history_str = "\n".join(history_lines) - - # Render template - template = Template(self.templates.MCTS_POLICY_TEMPLATE) - return template.render( - bos_token=self._get_bos_token(), - game_state=game_state, - valid_actions=actions_str, - history=history_str, - ) - - def build_sft_prompt( - self, - game_state: str, - valid_actions: List[str], - mcts_action: str, - mcts_value: float, - ) -> str: - """ - Build a prompt for supervised fine-tuning from MCTS policy. - - Args: - game_state: Current observation text - valid_actions: List of valid action strings - mcts_action: Action recommended by MCTS - mcts_value: Value estimate from MCTS - - Returns: - Formatted prompt string - """ - actions_str = "\n".join([f"{i+1}. {action}" for i, action in enumerate(valid_actions)]) - - template = Template(self.templates.SFT_FROM_MCTS_TEMPLATE) - return template.render( - bos_token=self._get_bos_token(), - game_state=game_state, - valid_actions=actions_str, - mcts_action=mcts_action, - mcts_value=f"{mcts_value:.3f}", - ) - - def build_rft_prompt( - self, - game_state: str, - valid_actions: List[str], - trajectory: List[Dict[str, Any]], - cumulative_reward: float, - ) -> str: - """ - Build a prompt for reward fine-tuning. - - Args: - game_state: Current observation text - valid_actions: List of valid action strings - trajectory: Recent trajectory with rewards - cumulative_reward: Total reward accumulated - - Returns: - Formatted prompt string - """ - actions_str = "\n".join([f"{i+1}. {action}" for i, action in enumerate(valid_actions)]) - - # Format trajectory with rewards - traj_lines = [] - for i, step in enumerate(trajectory[-5:]): - action = step.get('action', 'N/A') - reward = step.get('reward', 0) - traj_lines.append(f" Step {i+1}: Action: {action} → Reward: {reward:+.2f}") - trajectory_str = "\n".join(traj_lines) - - template = Template(self.templates.RFT_TEMPLATE) - return template.render( - bos_token=self._get_bos_token(), - game_state=game_state, - valid_actions=actions_str, - trajectory_with_rewards=trajectory_str, - cumulative_reward=f"{cumulative_reward:+.2f}", - ) - - def build_eval_prompt( - self, - game_state: str, - valid_actions: List[str], - ) -> str: - """ - Build a simple prompt for evaluation. - - Args: - game_state: Current observation text - valid_actions: List of valid action strings - - Returns: - Formatted prompt string - """ - actions_str = "\n".join([f"{i+1}. {action}" for i, action in enumerate(valid_actions)]) - - template = Template(self.templates.EVAL_TEMPLATE) - return template.render( - bos_token=self._get_bos_token(), - game_state=game_state, - valid_actions=actions_str, - ) - - -# ============================================================================== -# Utility Functions -# ============================================================================== - -def extract_action_from_llm_output(llm_output: str, valid_actions: List[str]) -> Optional[str]: - """ - Extract the action from LLM output with tags. - - Args: - llm_output: Full LLM response including and tags - valid_actions: List of valid action strings to match against - - Returns: - Extracted action string, or None if extraction fails - - Example: - >>> output = "I need light take lamp" - >>> extract_action_from_llm_output(output, ["take lamp", "go north"]) - "take lamp" - """ - import re - - # Pattern to extract content between and - pattern = r"\s*(.*?)\s*" - match = re.search(pattern, llm_output, re.DOTALL | re.IGNORECASE) - - if not match: - return None - - extracted = match.group(1).strip() - - # Try exact match first - if extracted in valid_actions: - return extracted - - # Try case-insensitive match - extracted_lower = extracted.lower() - for action in valid_actions: - if action.lower() == extracted_lower: - return action - - # Try fuzzy match (substring) - for action in valid_actions: - if extracted_lower in action.lower() or action.lower() in extracted_lower: - return action - - return None - - -# ============================================================================== -# Example Usage -# ============================================================================== - -if __name__ == "__main__": - print("="*80) - print("PriorZero Prompt Templates - Example Usage") - print("="*80) - - # Mock tokenizer - class MockTokenizer: - bos_token_id = 1 - def decode(self, ids): - return "" - - tokenizer = MockTokenizer() - builder = PriorZeroPromptBuilder(tokenizer) - - # Example game state - game_state = "You are standing in an open field west of a white house." - valid_actions = ["go north", "go south", "go east", "open mailbox", "take mailbox"] - history = [ - {"observation": "West of House", "action": "look", "reward": 0}, - {"observation": "You see a mailbox", "action": "examine mailbox", "reward": 0}, - ] - - print("\n1. MCTS Policy Prompt:") - print("-"*80) - prompt = builder.build_mcts_policy_prompt(game_state, valid_actions, history) - print(prompt) - - print("\n2. SFT Prompt:") - print("-"*80) - sft_prompt = builder.build_sft_prompt(game_state, valid_actions, "open mailbox", 0.75) - print(sft_prompt) - - print("\n3. RFT Prompt:") - print("-"*80) - trajectory = [ - {"action": "go east", "reward": 0}, - {"action": "open mailbox", "reward": 5}, - ] - rft_prompt = builder.build_rft_prompt(game_state, valid_actions, trajectory, 5.0) - print(rft_prompt) - - print("\n4. Action Extraction:") - print("-"*80) - llm_output = "The mailbox might contain something useful. open mailbox" - extracted = extract_action_from_llm_output(llm_output, valid_actions) - print(f"LLM Output: {llm_output}") - print(f"Extracted Action: {extracted}") - - print("\n" + "="*80) - print("✓ All prompt templates demonstrated successfully!") - print("="*80) diff --git a/zoo/jericho/priorzero/run_prior_ablation.sh b/zoo/jericho/priorzero/run_prior_ablation.sh new file mode 100644 index 000000000..eb355ae90 --- /dev/null +++ b/zoo/jericho/priorzero/run_prior_ablation.sh @@ -0,0 +1,241 @@ +#!/usr/bin/env bash +# ============================================================================= +# Ablation Study: VLM Prior on LunarLander +# Runs all parameter combinations and saves results to ablation_results.json +# +# Usage (on GPU worker): +# cd zoo/jericho/priorzero +# bash run_ablation.sh +# ============================================================================= +set -euo pipefail + +PYTHON="/mnt/shared-storage-user/puyuan/xiongjyu/envs/rft/bin/python3" +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" +EVAL_SCRIPT="${SCRIPT_DIR}/scripts/eval_vl_prior.py" +OUTPUT_DIR="${SCRIPT_DIR}/ablation_output" +MERGED_JSON="${SCRIPT_DIR}/ablation_results.json" + +# NUM_EPISODES=20 +NUM_EPISODES=2 + +SEED=0 +MAX_STEPS=1000 + +mkdir -p "$OUTPUT_DIR" + +echo "==============================================" +echo " Ablation Study: VLM Prior on LunarLander" +echo " Episodes per combo: ${NUM_EPISODES}" +echo " Output dir: ${OUTPUT_DIR}" +echo "==============================================" + +# Counter for tracking progress +TOTAL=0 +DONE=0 + +# --------------------------------------------------------------------------- +# Define all combinations +# --------------------------------------------------------------------------- +# Format: "tag|policy|prompt_style|vlm_image_mode|image_size" +COMBOS=( + # "random_baseline|random|concise|current_only|64" + "vlm_concise_current_64|vlm|concise|current_only|64" + "vlm_concise_current_256|vlm|concise|current_only|256" + "vlm_concise_first_and_current_64|vlm|concise|first_and_current|64" + "vlm_concise_first_and_current_256|vlm|concise|first_and_current|256" + "vlm_legacy_current_64|vlm|legacy|current_only|64" + "vlm_legacy_current_256|vlm|legacy|current_only|256" + "vlm_legacy_first_and_current_64|vlm|legacy|first_and_current|64" + "vlm_legacy_first_and_current_256|vlm|legacy|first_and_current|256" +) + +TOTAL=${#COMBOS[@]} + +# --------------------------------------------------------------------------- +# Run each combination +# --------------------------------------------------------------------------- +for combo in "${COMBOS[@]}"; do + IFS='|' read -r TAG POLICY PROMPT_STYLE IMAGE_MODE IMAGE_SIZE <<< "$combo" + DONE=$((DONE + 1)) + OUTFILE="${OUTPUT_DIR}/${TAG}.json" + + echo "" + echo "----------------------------------------------" + echo " [${DONE}/${TOTAL}] Running: ${TAG}" + echo " policy=${POLICY} prompt=${PROMPT_STYLE} img_mode=${IMAGE_MODE} res=${IMAGE_SIZE}" + echo "----------------------------------------------" + + if [ "$POLICY" == "random" ]; then + $PYTHON "$EVAL_SCRIPT" \ + --policies random \ + --num_episodes "$NUM_EPISODES" \ + --seed "$SEED" \ + --max_steps "$MAX_STEPS" \ + --image_size "$IMAGE_SIZE" \ + --output "$OUTFILE" + else + $PYTHON "$EVAL_SCRIPT" \ + --policies vlm \ + --num_episodes "$NUM_EPISODES" \ + --seed "$SEED" \ + --max_steps "$MAX_STEPS" \ + --image_size "$IMAGE_SIZE" \ + --prompt_style "$PROMPT_STYLE" \ + --vlm_image_mode "$IMAGE_MODE" \ + --output "$OUTFILE" + fi + + echo " >> Saved to ${OUTFILE}" +done + +# --------------------------------------------------------------------------- +# Merge all results into one JSON +# --------------------------------------------------------------------------- +echo "" +echo "==============================================" +echo " Merging results..." +echo "==============================================" + +$PYTHON -c " +import json, glob, os + +merged = {} +for f in sorted(glob.glob('${OUTPUT_DIR}/*.json')): + with open(f) as fh: + data = json.load(fh) + # Each file has {tag: summary_dict} + merged.update(data) + +# Add metadata about the combination parameters for easier analysis +combo_meta = { + 'random_baseline': {'policy': 'random', 'prompt_style': '-', 'image_mode': '-', 'image_size': 64}, + 'vlm_concise_current_64': {'policy': 'vlm', 'prompt_style': 'concise', 'image_mode': 'current_only', 'image_size': 64}, + 'vlm_concise_current_256': {'policy': 'vlm', 'prompt_style': 'concise', 'image_mode': 'current_only', 'image_size': 256}, + 'vlm_concise_first_and_current_64': {'policy': 'vlm', 'prompt_style': 'concise', 'image_mode': 'first_and_current', 'image_size': 64}, + 'vlm_concise_first_and_current_256': {'policy': 'vlm', 'prompt_style': 'concise', 'image_mode': 'first_and_current', 'image_size': 256}, + 'vlm_legacy_current_64': {'policy': 'vlm', 'prompt_style': 'legacy', 'image_mode': 'current_only', 'image_size': 64}, + 'vlm_legacy_current_256': {'policy': 'vlm', 'prompt_style': 'legacy', 'image_mode': 'current_only', 'image_size': 256}, + 'vlm_legacy_first_and_current_64': {'policy': 'vlm', 'prompt_style': 'legacy', 'image_mode': 'first_and_current', 'image_size': 64}, + 'vlm_legacy_first_and_current_256': {'policy': 'vlm', 'prompt_style': 'legacy', 'image_mode': 'first_and_current', 'image_size': 256}, +} + +# Enrich each result with combo metadata +for key in merged: + # Match by checking if key starts with any combo tag + for tag, meta in combo_meta.items(): + if key == tag or key.startswith(tag.replace(tag.split('_')[0] + '_', '', 1)): + merged[key]['combo_meta'] = meta + break + # Fallback: try to match the 'policy' field in the result + if 'combo_meta' not in merged[key]: + for tag, meta in combo_meta.items(): + if merged[key].get('policy', '') == tag or tag in merged[key].get('policy', ''): + merged[key]['combo_meta'] = meta + break + +output = { + 'experiment': 'VLM Prior Ablation on LunarLander-v2', + 'num_episodes': ${NUM_EPISODES}, + 'seed': ${SEED}, + 'results': merged, +} + +with open('${MERGED_JSON}', 'w') as f: + json.dump(output, f, indent=2) + +print(f'Merged {len(merged)} results -> ${MERGED_JSON}') +" + +# --------------------------------------------------------------------------- +# Print summary table +# --------------------------------------------------------------------------- +echo "" +echo "==============================================" +echo " Printing analysis table..." +echo "==============================================" + +$PYTHON -c " +import json + +with open('${MERGED_JSON}') as f: + data = json.load(f) + +results = data['results'] + +# Print Markdown table +print() +print('| # | Configuration | Policy | Prompt | Image Mode | Resolution | Mean Reward | Std | Min | Max | Avg Steps |') +print('|---|--------------|--------|--------|------------|------------|-------------|-----|-----|-----|-----------|') + +# Sort: random first, then by reward descending +items = sorted(results.items(), key=lambda x: (x[1].get('combo_meta', {}).get('policy', '') != 'random', -x[1]['reward_mean'])) + +for i, (tag, r) in enumerate(items, 1): + meta = r.get('combo_meta', {}) + policy = meta.get('policy', r.get('policy', '?')) + prompt = meta.get('prompt_style', '-') + img_mode = meta.get('image_mode', '-') + img_size = meta.get('image_size', '-') + res_str = f'{img_size}x{img_size}' if img_size != '-' else '-' + + print(f'| {i} | {tag:45s} | {policy:6s} | {prompt:7s} | {img_mode:18s} | {res_str:10s} | {r[\"reward_mean\"]:11.2f} | {r[\"reward_std\"]:5.2f} | {r[\"reward_min\"]:5.0f} | {r[\"reward_max\"]:5.0f} | {r[\"steps_mean\"]:9.0f} |') + +print() + +# Quick analysis +random_reward = None +best_vlm_tag = None +best_vlm_reward = -1e9 + +for tag, r in results.items(): + meta = r.get('combo_meta', {}) + if meta.get('policy') == 'random': + random_reward = r['reward_mean'] + elif r['reward_mean'] > best_vlm_reward: + best_vlm_reward = r['reward_mean'] + best_vlm_tag = tag + +print('=== Quick Analysis ===') +if random_reward is not None: + print(f'Random baseline: {random_reward:.2f}') +if best_vlm_tag: + print(f'Best VLM config: {best_vlm_tag} -> {best_vlm_reward:.2f}') + if random_reward is not None: + diff = best_vlm_reward - random_reward + print(f'Improvement over random: {diff:+.2f} ({diff/abs(random_reward)*100:+.1f}%)') + +# Dimension analysis +print() +print('=== Dimension-wise Analysis ===') + +def avg_reward(filter_fn): + vals = [r['reward_mean'] for t, r in results.items() if filter_fn(t, r)] + return sum(vals)/len(vals) if vals else float('nan') + +# Concise vs Legacy +concise_avg = avg_reward(lambda t, r: r.get('combo_meta', {}).get('prompt_style') == 'concise') +legacy_avg = avg_reward(lambda t, r: r.get('combo_meta', {}).get('prompt_style') == 'legacy') +print(f'Concise prompt avg reward: {concise_avg:.2f}') +print(f'Legacy prompt avg reward: {legacy_avg:.2f}') +print(f' -> Concise vs Legacy delta: {concise_avg - legacy_avg:+.2f}') + +# Current-only vs First+Current +current_avg = avg_reward(lambda t, r: r.get('combo_meta', {}).get('image_mode') == 'current_only') +first_cur_avg = avg_reward(lambda t, r: r.get('combo_meta', {}).get('image_mode') == 'first_and_current') +print(f'Current-only avg reward: {current_avg:.2f}') +print(f'First+Current avg reward: {first_cur_avg:.2f}') +print(f' -> First+Current delta: {first_cur_avg - current_avg:+.2f}') + +# 64 vs 256 +res64_avg = avg_reward(lambda t, r: r.get('combo_meta', {}).get('image_size') == 64 and r.get('combo_meta', {}).get('policy') == 'vlm') +res256_avg = avg_reward(lambda t, r: r.get('combo_meta', {}).get('image_size') == 256 and r.get('combo_meta', {}).get('policy') == 'vlm') +print(f'64x64 avg reward: {res64_avg:.2f}') +print(f'256x256 avg reward: {res256_avg:.2f}') +print(f' -> Upscale delta: {res256_avg - res64_avg:+.2f}') +" + +echo "" +echo "==============================================" +echo " Ablation study complete!" +echo " Full results: ${MERGED_JSON}" +echo "==============================================" diff --git a/zoo/jericho/priorzero/scripts/eval_vl_prior.py b/zoo/jericho/priorzero/scripts/eval_vl_prior.py new file mode 100644 index 000000000..ced06b7a0 --- /dev/null +++ b/zoo/jericho/priorzero/scripts/eval_vl_prior.py @@ -0,0 +1,345 @@ +#!/usr/bin/env python3 +""" +Evaluate VLM prior quality by running episodes with different policies: + - random: uniform random action selection + - vlm: VLM prior (greedy argmax from VL model output) + +Usage (on GPU worker): + cd zoo/jericho/priorzero + python scripts/eval_vl_prior.py --vl_model Qwen2.5-VL-7b --num_episodes 20 + python scripts/eval_vl_prior.py --vl_model Qwen2.5-VL-7b --num_episodes 20 --prompt_style legacy + python scripts/eval_vl_prior.py --vl_model Qwen2.5-VL-7b --num_episodes 20 --vlm_image_mode first_and_current + python scripts/eval_vl_prior.py --policies random # random-only baseline (no GPU needed) +""" +import argparse +import sys +import os +import time +import json +import glob +import numpy as np +from collections import defaultdict, deque +from pathlib import Path + + +# # --------------------------------------------------------------------------- +# # Fix NVIDIA driver visibility in containers (must run before torch import) +# # --------------------------------------------------------------------------- +# def _fix_nvidia_env(): +# """Auto-detect NVIDIA driver libs and force-load libcuda before torch init.""" +# import ctypes + +# # 1. Patch LD_LIBRARY_PATH for child processes / nvidia-smi +# candidate_lib_dirs = [ +# "/usr/local/nvidia/lib64", +# "/usr/local/nvidia/lib", +# "/usr/lib/x86_64-linux-gnu", +# "/usr/lib64", +# ] +# candidate_bin_dirs = [ +# "/usr/local/nvidia/bin", +# "/usr/local/cuda/bin", +# ] +# for pattern in ["/usr/**/libcuda.so.1", "/lib/**/libcuda.so.1"]: +# for p in glob.glob(pattern, recursive=True): +# d = os.path.dirname(p) +# if d not in candidate_lib_dirs: +# candidate_lib_dirs.append(d) + +# ld_path = os.environ.get("LD_LIBRARY_PATH", "") +# for d in candidate_lib_dirs: +# if os.path.isdir(d) and d not in ld_path: +# ld_path = d + ":" + ld_path +# os.environ["LD_LIBRARY_PATH"] = ld_path + +# path = os.environ.get("PATH", "") +# for d in candidate_bin_dirs: +# if os.path.isdir(d) and d not in path: +# path = d + ":" + path +# os.environ["PATH"] = path + +# # 2. Force-load libcuda.so.1 into the current process so torch can find it. +# # Setting LD_LIBRARY_PATH alone is too late — the dynamic linker only +# # reads it at process start. ctypes.CDLL loads it immediately. +# for d in candidate_lib_dirs: +# libcuda = os.path.join(d, "libcuda.so.1") +# if os.path.isfile(libcuda): +# try: +# ctypes.CDLL(libcuda) +# except OSError: +# continue +# break + +# _fix_nvidia_env() + +# # Now safe to check CUDA +import torch +if not torch.cuda.is_available(): + print("[WARN] torch.cuda.is_available() = False. VLM policy will fail.") + print(f" LD_LIBRARY_PATH = {os.environ.get('LD_LIBRARY_PATH', '(unset)')}") + print(f" Searching libcuda.so.1 ...") + found = glob.glob("/usr/**/libcuda.so*", recursive=True) + \ + glob.glob("/lib/**/libcuda.so*", recursive=True) + print(f" Found: {found or 'NONE — this node has no GPU driver'}") + print(" If running on a GPU node, check that the NVIDIA driver is mounted into the container.") +else: + print(f"[OK] CUDA available: {torch.cuda.get_device_name(0)}") + + +# ── ensure project root is importable ── +SCRIPT_DIR = Path(__file__).resolve().parent.parent # zoo/jericho/priorzero +sys.path.insert(0, str(SCRIPT_DIR)) +sys.path.insert(0, str(SCRIPT_DIR / "src")) +PROJECT_ROOT = SCRIPT_DIR.parent.parent.parent # LightZero root +sys.path.insert(0, str(PROJECT_ROOT)) + +# ── PLACEHOLDER_MORE_IMPORTS ── + + +# --------------------------------------------------------------------------- +# Environment wrapper (thin, no DI-engine dependency) +# --------------------------------------------------------------------------- +class LunarLanderImageWrapper: + """Minimal wrapper around gymnasium LunarLander with image obs.""" + + ACTION_NAMES = ["NOOP", "LEFT_ENGINE", "MAIN_ENGINE", "RIGHT_ENGINE"] + + def __init__(self, image_size: int = 64, seed: int = 0): + try: + import gymnasium as gym + except ImportError: + import gym + import cv2 + self._cv2 = cv2 + self._env = gym.make("LunarLander-v2", render_mode="rgb_array") + self._image_size = image_size + self._seed = seed + self._timestep = 0 + + def reset(self): + self._env.reset(seed=self._seed) + self._timestep = 0 + return self._render() + + def step(self, action_idx: int): + _, reward, terminated, truncated, info = self._env.step(action_idx) + self._timestep += 1 + done = terminated or truncated + obs = self._render() + return obs, reward, done, info + + def _render(self) -> np.ndarray: + frame = self._env.render() # (H, W, 3) uint8 + frame = self._cv2.resize(frame, (self._image_size, self._image_size), + interpolation=self._cv2.INTER_AREA) + # CHW float32 [0,1] — same as LunarLanderImageEnv + return np.transpose(frame, (2, 0, 1)).astype(np.float32) / 255.0 + + def close(self): + self._env.close() + + +# --------------------------------------------------------------------------- +# Policy: Random +# --------------------------------------------------------------------------- +class RandomPolicy: + name = "random" + + def select_action(self, obs, history, valid_actions): + idx = np.random.randint(len(valid_actions)) + return idx, valid_actions[idx] + + +# --------------------------------------------------------------------------- +# Policy: VLM Prior (greedy) +# --------------------------------------------------------------------------- +class VLMPolicy: + """Wraps VLPriorGenerator for greedy action selection.""" + + def __init__(self, prior_generator): + self.pg = prior_generator + self.name = "vlm" + + def select_action(self, obs, history, valid_actions): + result = self.pg.generate_prior( + observation=obs, + action_candidates=valid_actions, + history=history, + temperature=0.01, # near-greedy + ) + idx = int(np.argmax(result["action_probs"])) + return idx, valid_actions[idx] + + +# --------------------------------------------------------------------------- +# Episode runner +# --------------------------------------------------------------------------- +def run_episode(env, policy, history_maxlen: int = 3, max_steps: int = 1000): + """Run one episode, return (total_reward, steps, action_counts).""" + obs = env.reset() + history = deque(maxlen=history_maxlen) + total_reward = 0.0 + action_counts = defaultdict(int) + + for step in range(max_steps): + action_idx, action_name = policy.select_action( + obs, list(history), LunarLanderImageWrapper.ACTION_NAMES + ) + action_counts[action_name] += 1 + next_obs, reward, done, info = env.step(action_idx) + history.append((obs, action_name, float(reward), step)) + total_reward += reward + obs = next_obs + if done: + break + + return total_reward, step + 1, dict(action_counts) + + +# --------------------------------------------------------------------------- +# Main +# --------------------------------------------------------------------------- +def build_vl_policy(args): + """Build VLPriorGenerator from args. Requires GPU.""" + from vl_config import VL_MODEL_CONFIGS, GAME_DESCRIPTIONS + from vl_engine import VLLMVLEngine + from prior_generator import VLPriorGenerator + + model_cfg = VL_MODEL_CONFIGS[args.vl_model] + limit_mm = {"image": 4 if args.vlm_image_mode != "current_only" else 1} + print(f"Loading VL model: {args.vl_model} ({model_cfg['model_path']})") + + # Use high-level VLLMVLEngine with standalone=True (no DDP) + vl_engine = VLLMVLEngine( + model_name="qwen2.5-vl", + model_path=model_cfg["model_path"], + tensor_parallel_size=model_cfg["tensor_parallel_size"], + gpu_memory_utilization=model_cfg["gpu_memory_utilization"], + max_model_len=4096, + enable_sleep=False, + limit_mm_per_prompt=limit_mm, + standalone=True, + ) + + pg = VLPriorGenerator( + vl_engine=vl_engine, + model_name=model_cfg["model_path"], + use_cot=args.use_cot, + game_description=GAME_DESCRIPTIONS.get("LunarLander-v2", ""), + vlm_image_mode=args.vlm_image_mode, + prompt_style=args.prompt_style, + ) + return VLMPolicy(pg) + + +def main(): + parser = argparse.ArgumentParser(description="Evaluate VLM prior vs random on LunarLander") + parser.add_argument("--num_episodes", type=int, default=10) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument("--max_steps", type=int, default=1000) + parser.add_argument("--history_length", type=int, default=3) + parser.add_argument("--image_size", type=int, default=64) + # VLM settings + parser.add_argument("--vl_model", type=str, default="Qwen3-VL-8b") + # parser.add_argument("--vl_model", type=str, default="Qwen2.5-VL-7b") + parser.add_argument("--use_cot", action="store_true", default=True) + parser.add_argument("--no_cot", action="store_true") + parser.add_argument("--vlm_image_mode", type=str, default="current_only", + choices=["current_only", "first_and_current", "all_history"]) + parser.add_argument("--prompt_style", type=str, default="concise", + choices=["concise", "legacy"]) + # Which policies to run + parser.add_argument("--policies", type=str, nargs="+", default=["random", "vlm"], + choices=["random", "vlm"]) + parser.add_argument("--output", type=str, default=None, + help="Path to save JSON results (default: stdout only)") + args = parser.parse_args() + + if args.no_cot: + args.use_cot = False + + # Build policies + policies = [] + for p in args.policies: + if p == "random": + policies.append(RandomPolicy()) + elif p == "vlm": + if not torch.cuda.is_available(): + print("[SKIP] vlm policy requires GPU but CUDA is not available. Skipping.") + continue + policies.append(build_vl_policy(args)) + + if not policies: + print("[ERROR] No policies to evaluate. Exiting.") + sys.exit(1) + + # Run evaluation + all_results = {} + for policy in policies: + tag = f"{policy.name}" + if hasattr(policy, "pg"): + tag += f"_{args.prompt_style}_{args.vlm_image_mode}" + if args.use_cot: + tag += "_cot" + + print(f"\n{'='*60}") + print(f"Policy: {tag} | Episodes: {args.num_episodes}") + print(f"{'='*60}") + + rewards = [] + steps_list = [] + action_totals = defaultdict(int) + + for ep in range(args.num_episodes): + env = LunarLanderImageWrapper(image_size=args.image_size, + seed=args.seed + ep) + t0 = time.time() + ep_reward, ep_steps, ep_actions = run_episode( + env, policy, + history_maxlen=args.history_length, + max_steps=args.max_steps, + ) + elapsed = time.time() - t0 + env.close() + + rewards.append(ep_reward) + steps_list.append(ep_steps) + for k, v in ep_actions.items(): + action_totals[k] += v + + print(f" ep {ep:3d}: reward={ep_reward:8.2f} steps={ep_steps:4d} " + f"time={elapsed:.1f}s actions={dict(ep_actions)}") + + # Summary + r = np.array(rewards) + summary = { + "policy": tag, + "num_episodes": args.num_episodes, + "reward_mean": float(r.mean()), + "reward_std": float(r.std()), + "reward_min": float(r.min()), + "reward_max": float(r.max()), + "steps_mean": float(np.mean(steps_list)), + "action_distribution": dict(action_totals), + } + all_results[tag] = summary + + print(f"\n Summary: mean={r.mean():.2f} ± {r.std():.2f} " + f"min={r.min():.2f} max={r.max():.2f} " + f"avg_steps={np.mean(steps_list):.0f}") + + # Final comparison + print(f"\n{'='*60}") + print("COMPARISON") + print(f"{'='*60}") + for tag, s in all_results.items(): + print(f" {tag:40s} reward={s['reward_mean']:8.2f} ± {s['reward_std']:.2f}") + + if args.output: + with open(args.output, "w") as f: + json.dump(all_results, f, indent=2) + print(f"\nResults saved to {args.output}") + + +if __name__ == "__main__": + main() diff --git a/zoo/jericho/priorzero/scripts/run_priorzero.sh b/zoo/jericho/priorzero/scripts/run_priorzero.sh new file mode 100644 index 000000000..9a1f75d3f --- /dev/null +++ b/zoo/jericho/priorzero/scripts/run_priorzero.sh @@ -0,0 +1,33 @@ + +#!/bin/bash +set -x + +# 1. 训练环境参数 +CUDA_DEVICES="0" +NPROC_PER_NODE=1 +MASTER_PORT=24554 + +# 2. 程序相关参数 +ENV_ID="detective.z5" # "zork1.z5" "acorncourt.z5" "omniquest.z5" +LOG_DIR="./data_priorzero/run_logs" +LLM_MODEL="qwen2.5-3b" # "qwen2.5-3b" "qwen2.5-7b" +mkdir -p "${LOG_DIR}" + +CURRENT_TIME=$(date +"%Y%m%d_%H%M%S") +LOG_FILE="${LOG_DIR}/log_${ENV_ID}_${LLM_MODEL}_${CURRENT_TIME}.txt" + +# 3. 设置环境变量 +export CUDA_VISIBLE_DEVICES="${CUDA_DEVICES}" +export PYTHONFAULTHANDLER=1 +export TORCH_DISTRIBUTED_DEBUG=DETAIL +export NCCL_DEBUG=INFO + + +torchrun \ + --nproc_per_node="${NPROC_PER_NODE}" \ + --master-port="${MASTER_PORT}" \ + ./src/priorzero_entry_sync.py \ + --use_cot \ + --env_id "${ENV_ID}" \ + --model "${LLM_MODEL}" \ + 2>&1 | tee "${LOG_FILE}" \ No newline at end of file diff --git a/zoo/jericho/priorzero/scripts/run_priorzero_ddp.sh b/zoo/jericho/priorzero/scripts/run_priorzero_ddp.sh new file mode 100644 index 000000000..0fd833029 --- /dev/null +++ b/zoo/jericho/priorzero/scripts/run_priorzero_ddp.sh @@ -0,0 +1,43 @@ + +#!/bin/bash +set -x + +# 1. 训练环境参数 +CUDA_DEVICES="0,1,2,3" +NPROC_PER_NODE=4 +MASTER_PORT=24554 + +# 2. 程序相关参数 +ENV_ID="detective.z5" # "zork1.z5" "acorncourt.z5" "omniquest.z5" +LOG_DIR="./data_priorzero/run_logs" +LLM_MODEL="qwen2.5-3b" # "qwen2.5-3b" "qwen2.5-7b" +USE_COT=false # true / false +mkdir -p "${LOG_DIR}" + +CURRENT_TIME=$(date +"%Y%m%d_%H%M%S") +LOG_FILE="${LOG_DIR}/log_${ENV_ID}_${LLM_MODEL}_${CURRENT_TIME}.txt" + +# 3. 设置环境变量 +export CUDA_VISIBLE_DEVICES="${CUDA_DEVICES}" +export PYTHONFAULTHANDLER=1 +export TORCH_DISTRIBUTED_DEBUG=DETAIL +export NCCL_DEBUG=INFO + +if [ "${USE_COT}" = true ]; then + torchrun \ + --nproc_per_node="${NPROC_PER_NODE}" \ + --master-port="${MASTER_PORT}" \ + ./src/priorzero_entry_sync_ddp.py \ + --use_cot \ + --env_id "${ENV_ID}" \ + --model "${LLM_MODEL}" \ + 2>&1 | tee "${LOG_FILE}" +else + torchrun \ + --nproc_per_node="${NPROC_PER_NODE}" \ + --master-port="${MASTER_PORT}" \ + ./src/priorzero_entry_sync_ddp.py \ + --env_id "${ENV_ID}" \ + --model "${LLM_MODEL}" \ + 2>&1 | tee "${LOG_FILE}" +fi \ No newline at end of file diff --git a/zoo/jericho/priorzero/scripts/run_priorzero_vl_lunarlander.sh b/zoo/jericho/priorzero/scripts/run_priorzero_vl_lunarlander.sh new file mode 100644 index 000000000..c6c453802 --- /dev/null +++ b/zoo/jericho/priorzero/scripts/run_priorzero_vl_lunarlander.sh @@ -0,0 +1,93 @@ +#!/bin/bash +# PriorZero VL Training on LunarLander-v2 (Image Input) +# +# Usage: +# bash run_priorzero_vl_lunarlander.sh [NUM_GPUS] [VL_MODEL] [SEED] [EXTRA_ARGS...] +# +# Examples: +# bash run_priorzero_vl_lunarlander.sh 4 Qwen2.5-VL-7b 0 +# bash run_priorzero_vl_lunarlander.sh 2 Qwen3-VL-2b 42 +# bash run_priorzero_vl_lunarlander.sh 1 Qwen3-VL-2b 0 --quick_test +# bash run_priorzero_vl_lunarlander.sh 1 Qwen3-VL-2b 0 --quick_test --no_cot --no_vl_fixed --mcts_mode wm_logits +# bash run_priorzero_vl_lunarlander.sh 1 Qwen3-VL-2b 0 --cot_weight 0.05 + +set -euo pipefail + +# ===================== Configurable Parameters ===================== +NUM_GPUS=${1:-4} +VL_MODEL=${2:-"Qwen2.5-VL-3b"} +SEED=${3:-0} +EXTRA_ARGS="${@:4}" +# CUDA_DEVICES=${CUDA_DEVICES:-"0,1,2,3"} +CUDA_DEVICES=${CUDA_DEVICES:-"0,1"} +# MASTER_PORT=${MASTER_PORT:-29500} +MASTER_PORT=${MASTER_PORT:-29501} + +# =================================================================== + +# DDP / NCCL debugging environment variables +export CUDA_VISIBLE_DEVICES="${CUDA_DEVICES}" +export PYTHONFAULTHANDLER=1 +export TORCH_DISTRIBUTED_DEBUG=DETAIL +export NCCL_DEBUG=INFO + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +ENV_ID="LunarLander-v2" +TIMESTAMP="$(date +%y%m%d_%H%M%S)" + +# ---- Parse key flags from EXTRA_ARGS for log naming ---- +COT_TAG="cot" +VL_FIXED_TAG="vlFixed" +MCTS_MODE="llm_plus_wm_logits" +COT_WEIGHT="0.1" +IMG_MODE="current_only" +LOGPROB_MODE="approximate" + +for arg in ${EXTRA_ARGS}; do + case "${prev_arg:-}" in + --mcts_mode) MCTS_MODE="$arg" ;; + --cot_weight) COT_WEIGHT="$arg" ;; + --vlm_image_mode) IMG_MODE="$arg" ;; + --logprob_mode) LOGPROB_MODE="$arg" ;; + esac + case "$arg" in + --no_cot) COT_TAG="noCot" ;; + --no_vl_fixed) VL_FIXED_TAG="vlTrain" ;; + esac + prev_arg="$arg" +done + +if [ "${COT_TAG}" = "cot" ]; then + COT_TAG="cot${COT_WEIGHT}" +fi + +# Build structured log directory: logs//// +LOG_DIR="${SCRIPT_DIR}/logs/LunarLander/${VL_MODEL}/${VL_FIXED_TAG}/${COT_TAG}_mcts_${MCTS_MODE}_img_${IMG_MODE}" +mkdir -p "${LOG_DIR}" +LOG_FILE="${LOG_DIR}/seed${SEED}_gpu${NUM_GPUS}_${TIMESTAMP}.log" + +echo "========================================" +echo "PriorZero VL - LunarLander-v2 (Image)" +echo "========================================" +echo "GPUs: ${NUM_GPUS}" +echo "VL Model: ${VL_MODEL}" +echo "Seed: ${SEED}" +echo "Extra Args: ${EXTRA_ARGS}" +echo "CUDA: ${CUDA_DEVICES}" +echo "Master Port: ${MASTER_PORT}" +echo "Log File: ${LOG_FILE}" +echo "========================================" + +cd "${SCRIPT_DIR}" + +torchrun \ + --nproc_per_node "${NUM_GPUS}" \ + --master-port "${MASTER_PORT}" \ + priorzero_entry_unified.py \ + --input_type image \ + --env_id "${ENV_ID}" \ + --vl_model "${VL_MODEL}" \ + --seed "${SEED}" \ + --max_iter 1e6 \ + ${EXTRA_ARGS} \ + 2>&1 | tee "${LOG_FILE}" diff --git a/zoo/jericho/priorzero/src/game_segment_priorzero.py b/zoo/jericho/priorzero/src/game_segment_priorzero.py new file mode 100644 index 000000000..f75f12e10 --- /dev/null +++ b/zoo/jericho/priorzero/src/game_segment_priorzero.py @@ -0,0 +1,204 @@ +import numpy as np +from typing import Optional, List, Any +from lzero.mcts.buffer.game_segment import GameSegment as OriginalGameSegment + + +class GameSegment(OriginalGameSegment): + + def __init__( + self, + action_space, + game_segment_length: int = 200, + config: Optional[Any] = None, + task_id: Optional[int] = None + ): + super().__init__(action_space, game_segment_length, config, task_id) + + self.raw_obs_segment = [] # Raw text observations + self.history_obs_segment = [] + self.llm_prior_per_tok_segment = [] # LLM prior per token (for debugging) + self.cot_prefix_segment = [] # CoT prefixes for reuse (optimization) + self.llm_action_segment = [] # Actions selected by LLM + + def reset(self, init_observations: List[np.ndarray], init_raw_obs, init_history_obs) -> None: + """ + [PRIORZERO-MODIFIED] + Reset the segment with initial observations. + + Args: + init_observations: List of initial frame stack observations + init_raw_obs: Initial raw text observation + init_history_obs: Initial history observations + """ + super().reset(init_observations) + self.raw_obs_segment.clear() + self.history_obs_segment.clear() + self.llm_prior_per_tok_segment.clear() + self.cot_prefix_segment.clear() # Clear CoT prefix segment + self.llm_action_segment.clear() + + # 以下结果均是第 t 时刻的结果 + self.raw_obs_segment.append(init_raw_obs) + self.history_obs_segment.append(init_history_obs) + self.llm_prior_per_tok_segment.append(None) + self.cot_prefix_segment.append(None) + self.llm_action_segment.append(None) + + def append( + self, + action: int, + obs: np.ndarray, + reward: float, + action_mask: np.ndarray, + to_play: int, + timestep: int = 0, + chance: int = 0, + raw_obs_text: Optional[str] = None, + history_obs: Optional[List[str]] = None, + llm_prior_per_tok = None, + cot_prefix: Optional[str] = None, + llm_action: Optional[str] = None, + **kwargs + ) -> None: + + super().append(action, obs, reward, action_mask, to_play, timestep, chance) + self.raw_obs_segment.append(raw_obs_text) + self.history_obs_segment.append(history_obs) + self.llm_prior_per_tok_segment.append(llm_prior_per_tok) + self.cot_prefix_segment.append(cot_prefix) + self.llm_action_segment.append(llm_action) + + def store_search_stats(self, visit_counts: List, root_value: List) -> None: + super().store_search_stats(visit_counts, root_value) + + def game_segment_to_array(self) -> None: + super().game_segment_to_array() + + def pad_over( + self, next_segment_observations: List, next_segment_rewards: List, next_segment_actions: List, next_segment_root_values: List, + next_segment_child_visits: List, next_segment_improved_policy: List = None, next_chances: List = None, + next_segment_raw_obs: List = None, next_segment_history_obs: List = None, next_segment_llm_prior_per_tok: List = None, + next_segment_cot_prefix: List = None, next_segment_llm_action: List = None + ) -> None: + super().pad_over( + next_segment_observations, next_segment_rewards, next_segment_actions, next_segment_root_values, + next_segment_child_visits, next_segment_improved_policy, next_chances + ) + assert len(next_segment_raw_obs) <= self.num_unroll_steps + self.td_steps + assert len(next_segment_history_obs) <= self.num_unroll_steps + self.td_steps + assert len(next_segment_llm_prior_per_tok) <= self.num_unroll_steps + self.td_steps + assert len(next_segment_cot_prefix) <= self.num_unroll_steps + self.td_steps + assert len(next_segment_llm_action) <= self.num_unroll_steps + self.td_steps + + import copy + if len(next_segment_history_obs) > 0: + # Check if llm_prior_per_tok is dict (LLM text games) or array (VL Atari) + if next_segment_llm_prior_per_tok and isinstance(next_segment_llm_prior_per_tok[0], dict): + # LLM text games: validate consistency + assert self.raw_obs_segment[-1] == next_segment_llm_prior_per_tok[0]['current_obs'] + assert self.history_obs_segment[-1] == next_segment_llm_prior_per_tok[0]['history'] + assert self.history_obs_segment[-1][-1][1] == self.llm_action_segment[-1] + assert next_segment_history_obs[0][-1][1] == next_segment_llm_action[0] + # For VL Atari: llm_prior_per_tok is numpy array, skip validation + + for raw_obs in next_segment_raw_obs: + self.raw_obs_segment.append(copy.deepcopy(raw_obs)) + for history_obs in next_segment_history_obs: + self.history_obs_segment.append(copy.deepcopy(history_obs)) + for lp in next_segment_llm_prior_per_tok: + self.llm_prior_per_tok_segment.append(copy.deepcopy(lp)) + for action in next_segment_llm_action: + self.llm_action_segment.append(copy.deepcopy(action)) + + # Handle CoT prefix padding (optimization for CoT reuse) + if next_segment_cot_prefix is not None: + for cot_prefix in next_segment_cot_prefix: + self.cot_prefix_segment.append(copy.deepcopy(cot_prefix)) + + def get_unroll_raw_obs(self, timestep: int, num_unroll_steps: int = 0, padding: bool = False) -> np.ndarray: + """ + Overview: + Get an observation of the correct format: o[t, t + stack frames + num_unroll_steps]. + Arguments: + - timestep (int): The time step. + - num_unroll_steps (int): The extra length of the observation frames. + - padding (bool): If True, pad frames if (t + stack frames) is outside of the trajectory. + """ + stacked_raw_obs = self.raw_obs_segment[timestep:timestep + self.frame_stack_num + num_unroll_steps] + if padding: + pad_len = self.frame_stack_num + num_unroll_steps - len(stacked_raw_obs) + if pad_len > 0: + pad_frames = [stacked_raw_obs[-1] for _ in range(pad_len)] + stacked_raw_obs = stacked_raw_obs + pad_frames + return stacked_raw_obs + + def get_unroll_histroy_obs(self, timestep: int, num_unroll_steps: int = 0, padding: bool = False) -> np.ndarray: + """ + Overview: + Get an observation of the correct format: o[t, t + stack frames + num_unroll_steps]. + Arguments: + - timestep (int): The time step. + - num_unroll_steps (int): The extra length of the observation frames. + - padding (bool): If True, pad frames if (t + stack frames) is outside of the trajectory. + """ + stacked_histroy_obs = self.history_obs_segment[timestep:timestep + self.frame_stack_num + num_unroll_steps] + if padding: + pad_len = self.frame_stack_num + num_unroll_steps - len(stacked_histroy_obs) + if pad_len > 0: + pad_frames = [stacked_histroy_obs[-1] for _ in range(pad_len)] + stacked_histroy_obs = stacked_histroy_obs + pad_frames + return stacked_histroy_obs + + def get_unroll_llm_prior_per_tok(self, timestep: int, num_unroll_steps: int = 0, padding: bool = False) -> np.ndarray: + """ + Return LLM prior per token aligned with actions for unroll window. + """ + stacked_prior = list(self.llm_prior_per_tok_segment[timestep:timestep + self.frame_stack_num + num_unroll_steps]) + if padding: + pad_len = self.frame_stack_num + num_unroll_steps - len(stacked_prior) + if pad_len > 0: + pad_frames = [stacked_prior[-1] for _ in range(pad_len)] + stacked_prior = stacked_prior + pad_frames + return stacked_prior + + def get_unroll_cot_prefix(self, timestep: int, num_unroll_steps: int = 0, padding: bool = False) -> List[str]: + """ + Return CoT prefixes aligned with observations for unroll window (CoT reuse optimization). + + Args: + timestep: The time step + num_unroll_steps: The extra length of the CoT prefix frames + padding: If True, pad frames if outside of trajectory + + Returns: + List of CoT prefix strings + """ + stacked_cot_prefix = list(self.cot_prefix_segment[timestep:timestep + self.frame_stack_num +num_unroll_steps]) + if padding: + pad_len = self.frame_stack_num + num_unroll_steps - len(stacked_cot_prefix) + if pad_len > 0: + # Pad with empty strings or last prefix + pad_frames = [stacked_cot_prefix[-1] for _ in range(pad_len)] + stacked_cot_prefix = stacked_cot_prefix + pad_frames + return stacked_cot_prefix + + def get_unroll_llm_action(self, timestep: int, num_unroll_steps: int = 0, padding: bool = False) -> List[str]: + """ + Return LLM actions aligned with observations for unroll window. + + Args: + timestep: The time step + num_unroll_steps: The extra length of the CoT prefix frames + padding: If True, pad frames if outside of trajectory + + Returns: + List of LLM action strings + """ + stacked_llm_action = list(self.llm_action_segment[timestep:timestep + self.frame_stack_num + num_unroll_steps]) + if padding: + pad_len = self.frame_stack_num + num_unroll_steps - len(stacked_llm_action) + if pad_len > 0: + # Pad with empty strings or last action + pad_frames = [stacked_llm_action[-1] for _ in range(pad_len)] + stacked_llm_action = stacked_llm_action + pad_frames + return stacked_llm_action \ No newline at end of file diff --git a/zoo/jericho/priorzero/src/models/actor.py b/zoo/jericho/priorzero/src/models/actor.py new file mode 100644 index 000000000..d188fe9d0 --- /dev/null +++ b/zoo/jericho/priorzero/src/models/actor.py @@ -0,0 +1,720 @@ +from contextlib import contextmanager +from typing import Optional, Union, List, Dict +from collections import defaultdict +import os +import math +from tqdm import tqdm +import numpy as np +import deepspeed +from torch.optim import Optimizer +import torch +import torch.distributed as dist +import torch.nn as nn +from transformers import AutoModelForCausalLM, AutoModelForVision2Seq, AutoConfig, BitsAndBytesConfig +from peft import LoraConfig, PeftModel, TaskType, get_peft_model +from transformers import AutoModelForCausalLM, BitsAndBytesConfig +from transformers.integrations.deepspeed import HfDeepSpeedConfig +from transformers.trainer import get_scheduler + +from utils import compute_approx_kl, compute_entropy, masked_mean, torch_dist_barrier_and_cuda_sync, log_probs_from_logits + +def _normalize_vllm_weight_name(name: str) -> str: + if name.startswith("base_model.model."): + name = name[len("base_model.model."):] + name = name.replace(".base_layer.", ".") + return name + + +def _should_skip_vllm_sync_param(name: str) -> bool: + return any(marker in name for marker in ("lora_A", "lora_B", "lora_embedding_A", "lora_embedding_B")) + + +def _validate_vllm_sync_config(args, train_mode: str, vllm_engine) -> None: + if vllm_engine is None: + return + + ds_tensor_parallel_size = getattr(args, "ds_tensor_parallel_size", 1) + zero_stage = getattr(args, "zero_stage", 2) + + if ds_tensor_parallel_size != 1: + raise NotImplementedError( + "PolicyModel._deepspeed_broadcast currently supports only ds_tensor_parallel_size == 1. " + f"Got ds_tensor_parallel_size={ds_tensor_parallel_size}. " + "The active vLLM sync path does not safely handle DeepSpeed tensor parallel shards yet." + ) + + if zero_stage == 3 and train_mode == "lora": + raise NotImplementedError( + "PolicyModel._deepspeed_broadcast does not support train_mode='lora' with zero_stage=3. " + "This path needs adapter merge/unmerge together with ZeRO-3 sharded parameters, which is not " + "validated in the current implementation." + ) + +class Actor(nn.Module): + """ + Base class for Actor models in reinforcement learning. + + This class serves as a foundation for implementing various actor models, which are responsible for selecting actions based on the policy learned from the environment. + + Args: + pretrain_or_model (nn.Module): A pretrained model or a new model instance to be used as the actor. + attn_implementation (str, optional): Attention mechanism implementation to use. Defaults to "flash_attention_2". + bf16 (bool, optional): Enable bfloat16 precision for model computations. Defaults to True. + ds_config (dict, optional): Configuration for DeepSpeed, enabling model partitioning across multiple GPUs. Defaults to None. + device_map (dict, optional): Device mapping for loading the model onto specific devices. Defaults to None. + temperature (float, optional): Temperature for action selection. Defaults to 1.0. + """ + + def __init__( + self, + pretrain_or_model: str, + attn_implementation="flash_attention_2", + bf16=True, + ds_config=None, + device_map=None, + temperature=1.0, + train_mode_cfg=None, + **kwargs, + ) -> None: + super().__init__() + + self.temperature = temperature + self.pretrain_or_model = pretrain_or_model + self.train_mode_cfg = train_mode_cfg if train_mode_cfg is not None else {"mode": "full"} + self.train_mode = self.train_mode_cfg.get("mode", "full") + attn_impl = attn_implementation + + if ds_config is not None and ds_config["zero_optimization"]["stage"] == 3: + _ = HfDeepSpeedConfig(ds_config) + else: + _ = None + + # Detect if model is VL (Vision-Language) or LLM (Language Model) + config = AutoConfig.from_pretrained(pretrain_or_model, trust_remote_code=True) + is_vl = hasattr(config, 'vision_config') or 'VL' in config.__class__.__name__ + + if is_vl: + # Use AutoModelForVision2Seq for VL models (e.g., Qwen2.5-VL, Qwen3-VL) + self.model = AutoModelForVision2Seq.from_pretrained( + pretrain_or_model, + trust_remote_code=True, + attn_implementation=attn_impl, + torch_dtype=torch.bfloat16 if bf16 else "auto", + device_map=device_map, + ) + else: + # Use AutoModelForCausalLM for text-only LLM models + self.model = AutoModelForCausalLM.from_pretrained( + pretrain_or_model, + trust_remote_code=True, + attn_implementation=attn_impl, + torch_dtype=torch.bfloat16 if bf16 else "auto", + device_map=device_map, + ) + self.model.config.use_cache = False + + if self.train_mode == "lora": + self.model.enable_input_require_grads() + target_modules = self.train_mode_cfg.get("lora_target_modules") + target_modules = list(target_modules) if target_modules else None + lora_config = LoraConfig( + task_type=TaskType.CAUSAL_LM, + inference_mode=False, + r=self.train_mode_cfg.get("lora_r", 16), + lora_alpha=self.train_mode_cfg.get("lora_alpha", 32), + lora_dropout=self.train_mode_cfg.get("lora_dropout", 0.05), + bias=self.train_mode_cfg.get("lora_bias", "none"), + target_modules=target_modules, + ) + self.model = get_peft_model(self.model, lora_config) + elif self.train_mode != "full": + raise ValueError(f"Unsupported train_mode: {self.train_mode}") + + self.model.config.use_cache = False + + def forward( + self, + sequences: torch.LongTensor, + action_mask: Optional[torch.Tensor] = None, + attention_mask: Optional[torch.Tensor] = None, + return_output=False, + return_entropy=False, + ) -> torch.Tensor: + + foward_attention_mask = attention_mask + rolled_sequences = torch.roll(sequences, shifts=-1, dims=1) + position_ids = attention_mask.long().cumsum(-1) - 1 + position_ids.masked_fill_(attention_mask == 0, 1) + + output = self.model(sequences, attention_mask=foward_attention_mask, position_ids=position_ids) + + if return_entropy: + # Training path (micro-batch size 4): cast to fp32 for entropy + flash cross-entropy + assert return_output + output["logits"] = output["logits"].to(torch.float32) + entropy = compute_entropy(output["logits"]) + setattr(output, "entropy", entropy[:, :-1]) + + log_probs = log_probs_from_logits(output["logits"], rolled_sequences, temperature=self.temperature) + + log_probs = log_probs[:, :-1] + + action_log_probs = log_probs[:, -action_mask.shape[1] :] * action_mask.float() + return (action_log_probs, output) if return_output else action_log_probs + + def gradient_checkpointing_enable(self, gradient_checkpointing_kwargs={"use_reentrant": False}): + self.model.gradient_checkpointing_enable(gradient_checkpointing_kwargs=gradient_checkpointing_kwargs) + + def gradient_checkpointing_disable(self): + self.model.gradient_checkpointing_disable() + + def print_trainable_parameters(self): + if hasattr(self.model, "print_trainable_parameters"): + self.model.print_trainable_parameters() + +class ReferenceModel: + def __init__(self, strategy, pretrain): + self.strategy = strategy + model = Actor( + pretrain, + attn_implementation=strategy.args.attn_implementation, + bf16=strategy.args.bf16, + ds_config=strategy.get_ds_eval_config( + offload=False + ), + temperature=strategy.args.temperature, + ) + self.model = strategy.prepare(model, is_rlhf=True) + self.model.eval() + self.micro_train_batch_size = self.strategy.args.micro_train_batch_size + + @torch.no_grad() + def forward( + self, + sequences: torch.LongTensor, + action_mask: torch.Tensor, + attention_mask: torch.Tensor, + ) -> torch.Tensor: + """ + Return: action_log_probs [B, T_action] + """ + device = torch.cuda.current_device() + B = sequences.size(0) + outs = [] + chunk_size = self.micro_train_batch_size + + for i in range(0, B, chunk_size): + s = sequences[i : i + chunk_size].to(device) + am = action_mask[i : i + chunk_size].to(device) + attn = attention_mask[i : i + chunk_size].to(device) + + out = self.model( + s, + action_mask=am, + attention_mask=attn, + ) + outs.append(out) + return torch.cat(outs, dim=0) + +class BatchPPOTrainer: + def __init__( + self, + strategy, + actor, + actor_optim, + actor_scheduler, + micro_train_batch_size: int = 8, + vllm_engine = None + ): + self.strategy = strategy + self.args = strategy.args + + self.actor = actor + self.actor_optim = actor_optim + self.actor_scheduler = actor_scheduler + self.vllm_engine = vllm_engine + self.use_cuda_ipc = self.args.use_cuda_ipc + + self.micro_train_batch_size = micro_train_batch_size + from models.loss import PolicyLoss + self.policy_loss = PolicyLoss( + clip_eps_low=self.args.eps_clip_low_high[0], + clip_eps_high=self.args.eps_clip_low_high[1], + policy_loss_type=self.args.policy_loss_type, + enable_vllm_is_correction=self.args.enable_vllm_is_correction, + vllm_is_truncated_threshold=self.args.vllm_is_truncated_threshold, + use_cot=self.args.use_cot, + cot_weight=self.args.cot_weight, + use_mispo=self.args.use_mispo, + mispo_token_truncated_threshold=self.args.mispo_token_truncated_threshold, + mispo_traj_truncated_threshold=self.args.mispo_traj_truncated_threshold + ) + self.train_iter = 0 + + def train_batch(self, batch_data: Dict[str, torch.Tensor], kl_ctl: float, step_idx: int = 0) -> Dict[str, float]: + device = torch.cuda.current_device() + for k, v in batch_data.items(): + if torch.is_tensor(v): + batch_data[k] = v.to(device) + + all_samples_size = batch_data["input_ids"].size(0) + status_list = [] + pbar = tqdm( + range(0, all_samples_size, self.micro_train_batch_size), + desc=f"PPO batch step={step_idx}", + disable=not self.strategy.is_rank_0(), + ) + acc_grad_steps = self.strategy.accumulated_gradient + metrics_buffer = defaultdict(list) # 用于累积 micro_step 指标的缓冲区 + kl_early_stop_threshold = getattr(self.args, 'kl_early_stop_threshold', None) + kl_early_stopped = False + for micro_step, start_idx in enumerate(pbar): + end_idx = min(start_idx + self.micro_train_batch_size, all_samples_size) + micro_batch = { + 'input_ids': batch_data['input_ids'][start_idx:end_idx], + "attention_mask": batch_data['attention_mask'][start_idx:end_idx], + "action_mask": batch_data['action_mask'][start_idx:end_idx], + "advantages": batch_data['advantages'][start_idx:end_idx], + "old_action_log_probs": batch_data['old_action_log_probs'][start_idx:end_idx], + "log_status": batch_data['log_status'][start_idx:end_idx], + "rollout_action_logprob": batch_data['rollout_action_logprob'][start_idx:end_idx], + } + micro_batch['ref_action_log_probs'] = batch_data['ref_action_log_probs'][start_idx:end_idx] if batch_data['ref_action_log_probs'] is not None else None + action_log_probs, output = self.actor( + micro_batch['input_ids'], + micro_batch['action_mask'], + attention_mask=micro_batch['attention_mask'], + return_output=True, + return_entropy=True, + ) + actor_loss, clipfrac, clip_ratio, approx_kl, vllm_kl, mispo_token_mask, mispo_traj_mask = self.policy_loss( + input_ids=micro_batch['input_ids'], + log_probs=action_log_probs, + old_log_probs=micro_batch['old_action_log_probs'], + advantages=micro_batch['advantages'], + action_mask=micro_batch['action_mask'], + rollout_log_probs=micro_batch['rollout_action_logprob'] + ) + + if self.args.rft_kl_coef > 0 and micro_batch['ref_action_log_probs'] is not None: + kl = compute_approx_kl( + action_log_probs, + micro_batch['ref_action_log_probs'], + kl_estimator=self.args.kl_estimator + ) + kl_loss = masked_mean(kl, micro_batch["action_mask"]) + else: + kl_loss = torch.tensor(0.0, device=device) + + # KL early stopping: skip remaining micro-batches if ref_kl exceeds threshold + kl_loss_item_for_check = kl_loss.detach().float().item() + if kl_early_stop_threshold is not None and kl_early_stop_threshold > 0: + if kl_loss_item_for_check > kl_early_stop_threshold: + if self.strategy.is_rank_0() and not kl_early_stopped: + import logging + logging.getLogger("priorzero.train").warning( + f"[KL Early Stop] ref_kl={kl_loss_item_for_check:.4f} > threshold={kl_early_stop_threshold}, " + f"skipping gradient updates for remaining micro-batches at micro_step={micro_step}" + ) + kl_early_stopped = True + # Skip backward pass but still collect metrics for logging + entropy_loss = masked_mean(output.entropy[:, -micro_batch["action_mask"].shape[1] :], micro_batch["action_mask"]) + policy_loss_item = actor_loss.detach().float().item() + clipfrac_item = clipfrac.detach().float().item() + clip_ratio_item = clip_ratio.detach().float().item() + approx_kl_item = approx_kl.detach().float().item() + kl_loss_item = kl_loss_item_for_check + entropy_loss_item = entropy_loss.detach().float().item() + input_response_length_item = micro_batch["attention_mask"].sum().detach().float().item() / micro_batch["attention_mask"].shape[0] + response_length_item = micro_batch["action_mask"].sum().detach().float().item() / micro_batch["action_mask"].shape[0] + input_length_item = input_response_length_item - response_length_item + metrics_buffer["policy_loss"].append(policy_loss_item) + metrics_buffer["clipfrac"].append(clipfrac_item) + metrics_buffer["clip_ratio"].append(clip_ratio_item) + metrics_buffer["approx_kl"].append(approx_kl_item) + metrics_buffer["ref_kl"].append(kl_loss_item) + metrics_buffer["input_length"].append(input_length_item) + metrics_buffer["response_length"].append(response_length_item) + metrics_buffer['entropy'].append(entropy_loss_item) + log_status = micro_batch["log_status"] + other_status = {k: [item[k] for item in log_status] for k in log_status[0].keys()} + for k, v in other_status.items(): + metrics_buffer[k] = v + if ((micro_step + 1) % acc_grad_steps == 0) or ((micro_step + 1) == pbar.total): + self.train_iter += 1 + continue + + loss = actor_loss + kl_loss * float(kl_ctl.value) + + entropy_loss = masked_mean(output.entropy[:, -micro_batch["action_mask"].shape[1] :], micro_batch["action_mask"]) + if self.args.entropy_loss_coef != 0: + loss -= entropy_loss * self.args.entropy_loss_coef + + self.strategy.backward(loss, self.actor, self.actor_optim) + self.strategy.optimizer_step(self.actor_optim, self.actor, self.actor_scheduler, name="actor") + + policy_loss_item = actor_loss.detach().float().item() + clipfrac_item = clipfrac.detach().float().item() + clip_ratio_item = clip_ratio.detach().float().item() + approx_kl_item = approx_kl.detach().float().item() + kl_loss_item = kl_loss.detach().float().item() + input_response_length_item = micro_batch["attention_mask"].sum().detach().float().item() / micro_batch["attention_mask"].shape[0] + response_length_item = micro_batch["action_mask"].sum().detach().float().item() / micro_batch["action_mask"].shape[0] + input_length_item = input_response_length_item - response_length_item + entropy_loss_item = entropy_loss.detach().float().item() + total_loss_item = loss.detach().float().item() + + # PPO importance sampling ratio stats + with torch.no_grad(): + ratio = torch.exp(action_log_probs - micro_batch['old_action_log_probs']) + ratio_masked = (ratio * micro_batch['action_mask'].float()) + mask_sum = micro_batch['action_mask'].float().sum() + ratio_mean_item = (ratio_masked.sum() / mask_sum).item() if mask_sum > 0 else 1.0 + ratio_std_item = ((((ratio - ratio_mean_item) ** 2) * micro_batch['action_mask'].float()).sum() / mask_sum).sqrt().item() if mask_sum > 0 else 0.0 + + # Advantage stats for this micro-batch + adv = micro_batch['advantages'] + adv_mean_item = adv.mean().item() + adv_std_item = adv.std().item() if adv.numel() > 1 else 0.0 + + # Log prob means + log_prob_new_mean_item = masked_mean(action_log_probs, micro_batch['action_mask']).item() + log_prob_old_mean_item = masked_mean(micro_batch['old_action_log_probs'], micro_batch['action_mask']).item() + + kl_coef_item = float(kl_ctl.value) + + pbar.set_postfix({ + "policy_loss": policy_loss_item, + "clipfrac": clipfrac_item, + "approx_kl": approx_kl_item, + "iter": self.train_iter, + }) + + metrics_buffer["policy_loss"].append(policy_loss_item) + metrics_buffer["clipfrac"].append(clipfrac_item) + metrics_buffer["clip_ratio"].append(clip_ratio_item) + metrics_buffer["approx_kl"].append(approx_kl_item) + metrics_buffer["ref_kl"].append(kl_loss_item) + metrics_buffer["input_length"].append(input_length_item) + metrics_buffer["response_length"].append(response_length_item) + metrics_buffer['entropy'].append(entropy_loss_item) + metrics_buffer['total_loss'].append(total_loss_item) + metrics_buffer['ratio_mean'].append(ratio_mean_item) + metrics_buffer['ratio_std'].append(ratio_std_item) + metrics_buffer['advantage_mean'].append(adv_mean_item) + metrics_buffer['advantage_std'].append(adv_std_item) + metrics_buffer['log_prob_new_mean'].append(log_prob_new_mean_item) + metrics_buffer['log_prob_old_mean'].append(log_prob_old_mean_item) + metrics_buffer['kl_coef'].append(kl_coef_item) + if vllm_kl is not None: + metrics_buffer['vllm_kl'].append(vllm_kl.item()) + if mispo_token_mask is not None: + mispo_token_mask = mispo_token_mask * micro_batch["action_mask"] + metrics_buffer['mispo_token_ratio'].append((mispo_token_mask.sum() / micro_batch["action_mask"].sum()).item()) + if mispo_traj_mask is not None: + metrics_buffer['mispo_traj_ratio'].append((mispo_traj_mask.sum() / mispo_traj_mask.shape[0]).item()) + + log_status = micro_batch["log_status"] + other_status = {k: [item[k] for item in log_status] for k in log_status[0].keys()} + for k, v in other_status.items(): + metrics_buffer[k] = v + + if ((micro_step + 1) % acc_grad_steps == 0) or ((micro_step + 1) == pbar.total): + self.train_iter += 1 + status = { + "policy_loss": np.mean(metrics_buffer['policy_loss']), + "clipfrac": np.mean(metrics_buffer['clipfrac']), + "clip_ratio": np.mean(metrics_buffer['clip_ratio']), + "approx_kl": np.mean(metrics_buffer['approx_kl']), + "ref_kl": np.mean(metrics_buffer['ref_kl']), + "entropy": np.mean(metrics_buffer['entropy']), + + "iter": self.train_iter, + "lr": self.actor_scheduler.get_last_lr()[0], + "global_grad_norm": self.actor_optim._global_grad_norm, + + "input_length_max": np.max(metrics_buffer['input_length']), + "input_length_mean": np.mean(metrics_buffer['input_length']), + "input_length_min": np.min(metrics_buffer['input_length']), + + "response_length_max": np.max(metrics_buffer['response_length']), + "response_length_mean": np.mean(metrics_buffer['response_length']), + "response_length_min": np.min(metrics_buffer['response_length']), + + "value_advantage_max": np.max(metrics_buffer['value_advantage']), + "value_advantage_mean": np.mean(metrics_buffer['value_advantage']), + "value_advantage_min": np.min(metrics_buffer['value_advantage']), + + "total_loss": np.mean(metrics_buffer['total_loss']), + "ratio_mean": np.mean(metrics_buffer['ratio_mean']), + "ratio_std": np.mean(metrics_buffer['ratio_std']), + "advantage_mean": np.mean(metrics_buffer['advantage_mean']), + "advantage_std": np.mean(metrics_buffer['advantage_std']), + "log_prob_new_mean": np.mean(metrics_buffer['log_prob_new_mean']), + "log_prob_old_mean": np.mean(metrics_buffer['log_prob_old_mean']), + "kl_coef": np.mean(metrics_buffer['kl_coef']), + } + if "final_advantage" in metrics_buffer: + status["final_advantage_max"] = np.max(metrics_buffer['final_advantage']) + status["final_advantage_mean"] = np.mean(metrics_buffer['final_advantage']) + status["final_advantage_min"] = np.min(metrics_buffer['final_advantage']) + if "fmt_rewards" in metrics_buffer: + status["fmt_rewards"] = np.mean(metrics_buffer['fmt_rewards']) + if "vllm_kl" in metrics_buffer: + status["vllm_kl"] = np.mean(metrics_buffer['vllm_kl']) + + if "mispo_token_ratio" in metrics_buffer: + status["mispo_token_ratio"] = np.mean(metrics_buffer['mispo_token_ratio']) + if "mispo_traj_ratio" in metrics_buffer: + status["mispo_traj_ratio"] = np.mean(metrics_buffer['mispo_traj_ratio']) + if kl_early_stopped: + status["kl_early_stopped"] = 1.0 + metrics_buffer.clear() + + status = self.strategy.all_reduce(status) + status_list.append(status) + + return status_list + + def _deepspeed_broadcast(self): + _validate_vllm_sync_config(self.strategy.args, self.actor.train_mode, self.vllm_engine) + use_prefix_cache = getattr(self.strategy.args, "enable_prefix_caching", False) + if use_prefix_cache: + self.vllm_engine.reset_prefix_cache() + + torch.cuda.empty_cache() + model = self.actor.model.module + with self._merged_lora_adapter(model): + sync_params = list(self._iter_vllm_sync_params(model)) + count, num_params = 0, len(sync_params) + for name, param in sync_params: + count += 1 # empty_cache at last param + # For ZeRO-3, allgather sharded parameter and broadcast to all vllm engines by rank 0 + with deepspeed.zero.GatheredParameters([param], enabled=self.strategy.args.zero_stage == 3): + shape = param.shape if self.strategy.args.zero_stage != 3 else param.ds_shape + self.vllm_engine.update_weight(name, dtype=param.dtype, shape=shape, weight=param.data, empty_cache=(count == num_params)) + + def _broadcast_to_vllm(self): + use_prefix_cache = getattr(self.strategy.args, "enable_prefix_caching", False) + if use_prefix_cache and torch.distributed.get_rank() == 0: + self.vllm_engine.reset_prefix_cache() + + torch.cuda.empty_cache() + model = self.actor.model + count, num_params = 0, len(list(model.named_parameters())) + + def _broadcast_param(param, count, num_params): + if torch.distributed.get_rank() == 0: + shape = param.shape if self.strategy.args.zero_stage != 3 else param.ds_shape + self.vllm_engine.update_weight(name, dtype=param.dtype, shape=shape, empty_cache=count == num_params) + + self._model_update_group.broadcast(param.data, src=0, stream=torch.cuda.current_stream()) + + def _handle_cuda_ipc(param, count, num_params): + from torch.multiprocessing.reductions import reduce_tensor + + weight = param.data.clone() + ipc_handle = reduce_tensor(weight) + + from vllm_utils.vllm_engine import get_physical_gpu_id + ipc_handle = {get_physical_gpu_id(): ipc_handle} + ipc_handle_list = [None] * torch.distributed.get_world_size() + torch.distributed.all_gather_object(ipc_handle_list, ipc_handle) + + if torch.distributed.get_rank() == 0: + ipc_handles = {} + for d in ipc_handle_list: + ipc_handles.update(d) + + shape = param.shape if self.strategy.args.zero_stage != 3 else param.ds_shape + self.vllm_engine.update_weight_cuda_ipc( + name, + dtype=param.dtype, + shape=shape, + ipc_handles=ipc_handles, + empty_cache=count == num_params, + ) + + torch_dist_barrier_and_cuda_sync() + + for name, param in model.named_parameters(): + count += 1 # empty_cache at last param + + # broadcast + if not self.use_cuda_ipc: + # For ZeRO-3, allgather sharded parameter and broadcast to all vllm engines by rank 0 + if self.strategy.args.ds_tensor_parallel_size > 1: + with deepspeed.module_inject.layers.GatherReplacedLayerParams([param], model, enabled=True): + _broadcast_param(param, count, num_params) + else: + with deepspeed.zero.GatheredParameters([param], enabled=self.strategy.args.zero_stage == 3): + _broadcast_param(param, count, num_params) + else: + if self.strategy.args.ds_tensor_parallel_size > 1: + with deepspeed.module_inject.layers.GatherReplacedLayerParams([param], model, enabled=True): + _handle_cuda_ipc(param, count, num_params) + else: + with deepspeed.zero.GatheredParameters([param], enabled=self.strategy.args.zero_stage == 3): + _handle_cuda_ipc(param, count, num_params) + + torch.cuda.empty_cache() + torch_dist_barrier_and_cuda_sync() + + def _iter_vllm_sync_params(self, model): + for name, param in model.named_parameters(): + if _should_skip_vllm_sync_param(name): + continue + yield _normalize_vllm_weight_name(name), param + + @contextmanager + def _merged_lora_adapter(self, model): + if isinstance(model, PeftModel): + if not hasattr(model, "merge_adapter") or not hasattr(model, "unmerge_adapter"): + raise RuntimeError("Current PEFT version does not support merge_adapter/unmerge_adapter required for vLLM sync.") + model.merge_adapter() + try: + yield model + finally: + model.unmerge_adapter() + else: + yield model + + +class PolicyModel: + def __init__( + self, + strategy, + pretrain: str, + max_steps: Optional[int] = None, + vllm_engine=None, + ): + self.strategy = strategy + args = strategy.args + + self.vllm_engine = vllm_engine + self.max_steps = max_steps + + if getattr(args, "vllm_num_engines", 0) > 0: + if getattr(args, "vllm_sync_backend", "nccl") == "nccl": + os.environ["NCCL_CUMEM_ENABLE"] = "0" + + actor = Actor( + pretrain, + attn_implementation=args.attn_implementation, + bf16=args.bf16, + ds_config=strategy.get_ds_train_config(is_actor=True), + temperature=args.temperature, + train_mode_cfg=args.train_mode_dict, + ) + strategy.print(actor) + if args.train_mode_dict.mode == "lora": + actor.print_trainable_parameters() + + from transformers import AutoTokenizer + self.tokenizer = AutoTokenizer.from_pretrained( + pretrain, trust_remote_code=True, padding_side="left" + ) + if self.tokenizer.pad_token is None: + self.tokenizer.pad_token = self.tokenizer.eos_token + + actor_optim = strategy.create_optimizer( + actor, + lr=args.learning_rate, + betas=args.adam_betas, + weight_decay=args.weight_decay, + ) + + if max_steps is None: + max_steps = int(getattr(args, "max_steps", 1_000_000)) + + actor_scheduler = get_scheduler( + args.lr_scheduler, + actor_optim, + num_warmup_steps=math.ceil(max_steps * args.lr_warmup_ratio), + num_training_steps=max_steps, + scheduler_specific_kwargs={"min_lr": args.learning_rate * 0.1}, + ) + + if args.gradient_checkpointing: + actor.gradient_checkpointing_enable( + gradient_checkpointing_kwargs={"use_reentrant": args.gradient_checkpointing_use_reentrant} + ) + + self.actor, self.actor_optim, self.actor_scheduler = strategy.prepare( + (actor, actor_optim, actor_scheduler), + is_rlhf=True, + ) + + if strategy.args.deepspeed_enable_sleep: + from strategy.deepspeed import offload_deepspeed_states + offload_deepspeed_states(self.actor.model) + + self.trainer = BatchPPOTrainer( + strategy, + self.actor, + actor_optim=self.actor_optim, + actor_scheduler=self.actor_scheduler, + micro_train_batch_size=args.micro_train_batch_size, + vllm_engine = vllm_engine, + ) + self.micro_train_batch_size = self.strategy.args.micro_train_batch_size + + def fit(self, batch_data, kl_ctl: float = 0.0): + torch.cuda.empty_cache() + self.actor.train() + status = self.trainer.train_batch(batch_data, kl_ctl) + torch.cuda.empty_cache() + torch.cuda.synchronize() + return status + + @torch.no_grad() + def forward( + self, + sequences: torch.LongTensor, + action_mask: torch.Tensor, + attention_mask: torch.Tensor, + ) -> torch.Tensor: + """ + Return: action_log_probs [B, T_action] + """ + self.actor.eval() + device = torch.cuda.current_device() + B = sequences.size(0) + + outs = [] + chunk_size = self.micro_train_batch_size + + for i in range(0, B, chunk_size): + s = sequences[i : i + chunk_size].to(device) + am = action_mask[i : i + chunk_size].to(device) + attn = attention_mask[i : i + chunk_size].to(device) + out = self.actor( + s, + action_mask=am, + attention_mask=attn, + ) + outs.append(out) + return torch.cat(outs, dim=0) + + def broadcast_to_vllm(self): + # self.trainer._broadcast_to_vllm() + self.trainer._deepspeed_broadcast() + + def save_model(self): + args = self.strategy.args + self.strategy.save_model( + self.actor, + self.tokenizer, + args.save_path, + ) + @property + def train_iter(self): + return self.trainer.train_iter + + def reload_states(self): + from strategy.deepspeed import reload_deepspeed_states + reload_deepspeed_states(self.actor.model) + + def offload_states(self): + from strategy.deepspeed import offload_deepspeed_states + offload_deepspeed_states(self.actor.model) \ No newline at end of file diff --git a/zoo/jericho/priorzero/src/models/loss.py b/zoo/jericho/priorzero/src/models/loss.py new file mode 100644 index 000000000..d33f543cd --- /dev/null +++ b/zoo/jericho/priorzero/src/models/loss.py @@ -0,0 +1,159 @@ +from typing import Optional, Tuple + +import torch +import torch.distributed as dist +import torch.nn as nn +import torch.nn.functional as F + +from utils import masked_mean + +class PolicyLoss(nn.Module): + """ + Policy Loss for PPO + """ + + def __init__( + self, + clip_eps_low: float = 0.2, + clip_eps_high: float = 0.2, + dual_clip: float = None, + token_level_loss: bool = True, + policy_loss_type: str = "ppo", + enable_vllm_is_correction: bool = False, + vllm_is_truncated_threshold: list = None, + use_icepop: bool = False, + use_cot: bool = False, + use_mispo: bool = False, + cot_weight: Optional[float] = None, + mispo_token_truncated_threshold = None, + mispo_traj_truncated_threshold = None, + + ) -> None: + super().__init__() + self.clip_eps_low = clip_eps_low + self.clip_eps_high = clip_eps_high + self.token_level_loss = token_level_loss + self.dual_clip = dual_clip + self.policy_loss_type = policy_loss_type + self.enable_vllm_is_correction = enable_vllm_is_correction + self.vllm_is_truncated_threshold = vllm_is_truncated_threshold + self.use_icepop = use_icepop + + self.use_cot = use_cot + self.cot_weight = cot_weight + self.use_mispo = use_mispo + self.mispo_token_truncated_threshold = mispo_token_truncated_threshold + self.mispo_traj_truncated_threshold = mispo_traj_truncated_threshold + + # GSPO requires sequence-level loss + if policy_loss_type == "gspo": + self.token_level_loss = False + + # Dual-clip PPO: https://arxiv.org/pdf/1912.09729 + if dual_clip is not None: + assert dual_clip > 1.0, f"dual_clip must be > 1.0, got {dual_clip}" + + def forward( + self, + input_ids: torch.LongTensor, + log_probs: torch.Tensor, + old_log_probs: torch.Tensor, + advantages: torch.Tensor, + action_mask: Optional[torch.Tensor] = None, + rollout_log_probs: Optional[torch.Tensor] = None, + ) -> torch.Tensor: + if self.policy_loss_type == "ppo": + log_ratio = log_probs - old_log_probs + ratio = log_ratio.exp() + elif self.policy_loss_type == "gspo": + # GSPO: https://arxiv.org/pdf/2507.18071 + if self.enable_vllm_is_correction: + log_ratio = log_probs - rollout_log_probs + else: + log_ratio = log_probs - old_log_probs + ratio = (log_ratio * action_mask).sum(dim=-1) / action_mask.sum(dim=-1) + ratio = ratio.exp().unsqueeze(-1) * action_mask + else: + raise ValueError(f"Invalid policy loss type: {self.policy_loss_type}") + if advantages.dim() == 1: + advantages = advantages.unsqueeze(-1) + + surr1 = ratio * advantages + surr2 = ratio.clamp(1 - self.clip_eps_low, 1 + self.clip_eps_high) * advantages + + if self.dual_clip is None: + # Standard PPO + loss = -torch.min(surr1, surr2) + else: + # Standard PPO clipping + clip1 = torch.min(surr1, surr2) + # Dual-clip: additional lower bound for negative advantages + clip2 = torch.max(clip1, self.dual_clip * advantages) + # Apply dual-clip: use clip2 for negative advantages, clip1 for positive advantages + loss = -torch.where(advantages < 0, clip2, clip1) + + # Your Efficient RL Framework Secretly Brings You Off-Policy RL Training: https://fengyao.notion.site/off-policy-rl + vllm_kl = None + token_mask = None + traj_mask = None + effective_mask = action_mask + if self.enable_vllm_is_correction and self.policy_loss_type == "ppo": + low_threshold, high_threshold = self.vllm_is_truncated_threshold + if self.use_mispo: + token_low, token_high = self.mispo_token_truncated_threshold + traj_low, traj_high = self.mispo_traj_truncated_threshold + token_ratio = torch.exp(old_log_probs - rollout_log_probs).detach() + token_mask = ((token_ratio >= token_low) & (token_ratio <= token_high)).float() + traj_log_ratio = masked_mean( + old_log_probs - rollout_log_probs, + action_mask, + dim=-1, + ) + traj_ratio = torch.exp(traj_log_ratio).detach() + traj_mask = ((traj_ratio >= traj_low) & (traj_ratio <= traj_high)).float().unsqueeze(-1) + mispo_mask = token_mask * traj_mask * action_mask + loss = loss * mispo_mask + effective_mask = mispo_mask + if effective_mask.sum().item() == 0: + effective_mask = action_mask + + elif self.use_icepop: + # ICEPOP: set coefficients outside the interval to 0 + vllm_is = torch.exp(old_log_probs - rollout_log_probs).detach() + mask = (vllm_is >= low_threshold) & (vllm_is <= high_threshold) + vllm_is = vllm_is * mask + loss = vllm_is * loss + else: + # Standard clamp with low and high thresholds + vllm_is = ( + torch.exp(old_log_probs - rollout_log_probs).clamp(min=low_threshold, max=high_threshold).detach() + ) + loss = vllm_is * loss + vllm_kl = masked_mean(rollout_log_probs - old_log_probs, effective_mask, dim=None) + + ###### 对 cot 前缀加权重 + if self.use_cot and self.cot_weight is not None: + output_ids = input_ids[:, -action_mask.shape[1]:] + is_split = (output_ids == 2512) & action_mask.bool() + token_weights = torch.ones_like(loss) + pos = torch.arange(action_mask.shape[1], device=input_ids.device).unsqueeze(0) + last_split_pos = torch.where(is_split, pos, torch.full_like(pos, -1)).max(dim=1, keepdim=True).values + token_weights = torch.where( + (pos < last_split_pos) & action_mask.bool(), # 若想包含 2512 本身就改成 <= + torch.full_like(token_weights, self.cot_weight), + token_weights, + ) + loss = loss * token_weights + + loss = ( + masked_mean(loss, effective_mask, dim=None) + if self.token_level_loss + else masked_mean(loss, effective_mask, dim=-1).mean() + ) + + clipped = ratio.gt(1 + self.clip_eps_high) | ratio.lt(1 - self.clip_eps_low) + clipfrac = masked_mean(clipped, effective_mask, dim=None) + + clip_ratio = masked_mean(torch.lt(surr2, surr1).float(), effective_mask, dim=None) + approx_kl = masked_mean(-log_ratio.detach(), effective_mask, dim=None) + return loss, clipfrac, clip_ratio, approx_kl, vllm_kl, token_mask, traj_mask \ No newline at end of file diff --git a/zoo/jericho/priorzero/src/models/stability_optimizer.py b/zoo/jericho/priorzero/src/models/stability_optimizer.py new file mode 100644 index 000000000..789ef0880 --- /dev/null +++ b/zoo/jericho/priorzero/src/models/stability_optimizer.py @@ -0,0 +1,156 @@ +import logging +from collections import deque +from typing import Dict, Optional, Tuple, Union + +import numpy as np +import torch + + +class AdaptiveValueNormalizer: + """ + 作用:把 value/return/advantage 变成稳定尺度(近似零均值、单位方差),并支持 soft(log-sym)/hard(percentile) 抑制极端值。 + 核心:batch 统计(只看当前) + EMA 运行统计(全局追踪非平稳) + 可选裁剪/压缩。 + """ + + def __init__( + self, + init_momentum: float = 0.9, + final_momentum: float = 0.99, + warmup_steps: int = 100, + clip_method: str = "soft", # "soft" | "hard" | "none" + clip_percentile: float = 0.95, # hard clip 中间保留比例,如 0.95 => 保留 [2.5%, 97.5%] + min_std: float = 1e-6, + hard_clip_start_updates: int = 10, # hard clip 前几次不启用 + history_size: int = 1000, + ): + self.init_momentum = init_momentum + self.final_momentum = final_momentum + self.warmup_steps = warmup_steps + self.clip_method = clip_method + self.clip_percentile = clip_percentile + self.min_std = min_std + self.hard_clip_start_updates = hard_clip_start_updates + + self.running_mean = 0.0 + self.running_std = 1.0 + self.update_count = 0 + + self.value_history = deque(maxlen=history_size) + + def _momentum(self) -> float: + if self.update_count >= self.warmup_steps: + return self.final_momentum + p = self.update_count / max(self.warmup_steps, 1) + return self.init_momentum + (self.final_momentum - self.init_momentum) * p + + @staticmethod + def _log_sym(x: torch.Tensor) -> Tuple[torch.Tensor, int]: + # f(x)=sign(x)*log(1+|x|) + significant = int((x.abs() > 10).sum()) + y = torch.sign(x) * torch.log1p(torch.abs(x)) + return y, significant + + def _hard_percentile_clip(self, x: torch.Tensor) -> Tuple[torch.Tensor, int]: + if self.update_count < self.hard_clip_start_updates: + return x, 0 + q = self.clip_percentile + lo = (1 - q) / 2 + hi = 1 - lo + + xf = x.flatten() + lb = torch.quantile(xf, lo) + ub = torch.quantile(xf, hi) + y = torch.clamp(x, lb, ub) + + clipped = int((y != x).sum()) + return y, clipped + + def _batch_mean_std(self, x: torch.Tensor) -> Tuple[float, float]: + xf = x.flatten() + n = xf.numel() + if n == 0: + return 0.0, 1.0 + if n == 1: + mean = float(xf.item()) + return mean, self.min_std + + xf64 = xf.to(torch.float64) + mean = float(xf64.mean().item()) + var = float(xf64.var(unbiased=True).item()) + std = max(var ** 0.5, self.min_std) + return mean, std + + def normalize( + self, + values: torch.Tensor, + clip_values: bool = True, + return_stats: bool = False, + ) -> Union[torch.Tensor, Tuple[torch.Tensor, Dict]]: + x = values.detach() + + clipped_count = 0 + if clip_values: + if self.clip_method == "soft": + x, clipped_count = self._log_sym(x) + elif self.clip_method == "hard": + x, clipped_count = self._hard_percentile_clip(x) + else: + raise ValueError(f"Unknown clip_method: {self.clip_method}") + + batch_mean, batch_std = self._batch_mean_std(x) + + m = self._momentum() + if self.update_count == 0: + self.running_mean = batch_mean + self.running_std = batch_std + else: + self.running_mean = m * self.running_mean + (1 - m) * batch_mean + self.running_std = m * self.running_std + (1 - m) * batch_std + + self.update_count += 1 + self.value_history.extend(x.flatten().float().cpu().tolist()) + + + y = (x.to(values.dtype) - self.running_mean) / (self.running_std + self.min_std) + + if not return_stats: + return y + + stats = { + "batch_mean": batch_mean, + "batch_std": batch_std, + "running_mean": self.running_mean, + "running_std": self.running_std, + "momentum": m, + "clip_method": self.clip_method, + "clipped_count": clipped_count, + "total_count": int(x.numel()), + } + return y, stats + + def summary(self) -> Dict: + if self.update_count == 0: + return {} + recent = list(self.value_history)[-min(100, len(self.value_history)) :] + return { + "total_updates": self.update_count, + "current_mean": float(self.running_mean), + "current_std": float(self.running_std), + "recent_mean": float(np.mean(recent)) if recent else 0.0, + "recent_std": float(np.std(recent)) if recent else 1.0, + "recent_min": float(np.min(recent)) if recent else 0.0, + "recent_max": float(np.max(recent)) if recent else 0.0, + "clip_method": self.clip_method, + } + + def clear(self): + """ + Reset all running statistics so the normalizer behaves like a fresh instance. + Useful when starting a new experiment or episode. + """ + self.running_mean = 0.0 + self.running_std = 1.0 + self.update_count = 0 + + self.value_history.clear() + diff --git a/zoo/jericho/priorzero/src/priorzero_collector.py b/zoo/jericho/priorzero/src/priorzero_collector.py new file mode 100644 index 000000000..4ae4da53f --- /dev/null +++ b/zoo/jericho/priorzero/src/priorzero_collector.py @@ -0,0 +1,733 @@ +import asyncio +import logging +import sys +import time + +from collections import deque, defaultdict +from pathlib import Path +from typing import Optional, Any, List, Dict, Tuple + +import numpy as np +import torch +import torch.distributed as dist +from ding.envs import BaseEnvManager +from ding.torch_utils import to_ndarray +from ding.utils import build_logger, EasyTimer, SERIAL_COLLECTOR_REGISTRY, allreduce_data +from vllm import SamplingParams +import os +import math + +# Import from local LightZero +from lzero.worker.muzero_segment_collector import MuZeroSegmentCollector as OriginalCollector +from lzero.mcts.utils import prepare_observation +from game_segment_priorzero import GameSegment + +# ============================================================================== +# Helper Functions +# ============================================================================== + +def extract_raw_obs_text(obs_dict: Dict[str, Any]) -> str: + """ + Extract text observation from environment observation dictionary. + + Args: + obs_dict: Observation dictionary from environment + + Returns: + text_obs: Text observation string + """ + # [PRIORZERO-FIX] Try to get 'raw_obs_text' field first (Jericho env adds this) + if 'raw_obs_text' in obs_dict: + return str(obs_dict['raw_obs_text']) + + # Try to get 'raw_obs' field (alternative naming) + if 'raw_obs' in obs_dict: + return str(obs_dict['raw_obs']) + + # Try to get 'text' field + if 'text' in obs_dict: + return str(obs_dict['text']) + + # Try to get 'observation_str' field (Jericho env provides this in save_replay mode) + if 'observation_str' in obs_dict: + return str(obs_dict['observation_str']) + + # Try to get 'observation' and check if it's text + if 'observation' in obs_dict: + obs = obs_dict['observation'] + if isinstance(obs, str): + return obs + elif isinstance(obs, (list, np.ndarray)): + # If observation is already processed (e.g., embeddings), cannot extract text + # Return a placeholder + return f"[Observation vector of shape {np.array(obs).shape}]" + + # Fallback: return str representation + return str(obs_dict) + +# ============================================================================== +# PriorZero Collector Class +# ============================================================================== + +@SERIAL_COLLECTOR_REGISTRY.register('priorzero_segment', force_overwrite=True) +class PriorZeroCollector(OriginalCollector): + """ + [PRIORZERO-MODIFIED] + + Features: + - History buffer for each environment (sliding window) + - Robust error handling with retries + - Detailed logging of LLM prior statistics + """ + + def __init__( + self, + policy_config: Dict, + llm_config: Dict, + data_processor = None, + prof = None, + **kwargs + ): + """ + Initialize PriorZeroCollector. + + Args: + vllm_engine + policy_config: Policy configuration + llm_config: llm configuration + **kwargs: Additional arguments for parent class + """ + kwargs['policy_config'] = policy_config + + super().__init__(**kwargs) + + self.data_processor = data_processor + self.prof = prof + self.llm_cfg = llm_config + + self.history_buffers = defaultdict( + lambda: deque(maxlen=self.llm_cfg.history_length) + ) + self.llm_prior_temperature = llm_config.llm_prior_temperature + + self._logger.info(f"[RANK {self._rank}] ✓ PriorZeroCollector initialized with vLLM engine") + self._logger.info(f"[RANK {self._rank}] - History length: {self.llm_cfg.history_length}") + self._logger.info(f"[RANK {self._rank}] - Generate max length: {self.llm_cfg.generate_max_len}") + + def _should_continue_collect(self, local_done: bool) -> bool: + # Mirror of PriorZeroEvaluator._should_continue_eval. With vLLM TP > 1 + # spanning DDP ranks, exiting collect() while a TP partner is still + # inside vllm.generate causes the partner to deadlock at the next + # _sync_prompts_for_tp. all_reduce(MAX) a 0/1 flag and continue while + # ANY rank still needs work. No-op for TP=1 / single-process. + tp_size = getattr(self.llm_cfg, 'vllm_tensor_parallel_size', 1) + if dist.is_initialized() and dist.get_world_size() > 1 and tp_size > 1: + flag = torch.tensor([0 if local_done else 1], dtype=torch.long, + device=torch.cuda.current_device()) + dist.all_reduce(flag, op=dist.ReduceOp.MAX) + return flag.item() > 0 + return not local_done + + def pad_and_save_last_trajectory( + self, i: int, last_game_segments: List[GameSegment], last_game_priorities: List[np.ndarray], + game_segments: List[GameSegment], done: np.ndarray + ) -> None: + beg_index = self.policy_config.model.frame_stack_num + end_index = beg_index + self.policy_config.num_unroll_steps + self.policy_config.td_steps + + pad_obs_lst = game_segments[i].obs_segment[beg_index:end_index] + pad_raw_obs_lst = game_segments[i].raw_obs_segment[beg_index:end_index] + pad_history_obs_lst = game_segments[i].history_obs_segment[beg_index:end_index] + pad_llm_prior_per_tok_lst = game_segments[i].llm_prior_per_tok_segment[beg_index:end_index] + pad_cot_prefix_lst = game_segments[i].cot_prefix_segment[beg_index:end_index] # CoT reuse + pad_llm_action_lst = game_segments[i].llm_action_segment[beg_index:end_index] + + # NOTE: Specific padding logic for UniZero. + pad_action_lst = game_segments[i].action_segment[:self.policy_config.num_unroll_steps + self.policy_config.td_steps] + pad_child_visits_lst = game_segments[i].child_visit_segment[:self.policy_config.num_unroll_steps + self.policy_config.td_steps] + + beg_index = 0 + end_index = beg_index + self.unroll_plus_td_steps - 1 + pad_reward_lst = game_segments[i].reward_segment[beg_index:end_index] + + if self.policy_config.use_ture_chance_label_in_chance_encoder: + chance_lst = game_segments[i].chance_segment[beg_index:end_index] + + beg_index = 0 + end_index = beg_index + self.unroll_plus_td_steps + pad_root_values_lst = game_segments[i].root_value_segment[beg_index:end_index] + + if self.policy_config.gumbel_algo: + pad_improved_policy_prob = game_segments[i].improved_policy_probs[beg_index:end_index] + + # Pad and finalize the last game segment. + if self.policy_config.gumbel_algo: + last_game_segments[i].pad_over( + pad_obs_lst, pad_reward_lst, pad_action_lst, pad_root_values_lst, pad_child_visits_lst, + next_segment_improved_policy=pad_improved_policy_prob, + next_segment_cot_prefix=pad_cot_prefix_lst, # CoT reuse + next_segment_llm_action=pad_llm_action_lst + ) + else: + if self.policy_config.use_ture_chance_label_in_chance_encoder: + last_game_segments[i].pad_over( + pad_obs_lst, pad_reward_lst, pad_action_lst, pad_root_values_lst, pad_child_visits_lst, + next_chances=chance_lst, next_segment_raw_obs=pad_raw_obs_lst, + next_segment_history_obs=pad_history_obs_lst, next_segment_llm_prior_per_tok=pad_llm_prior_per_tok_lst, + next_segment_cot_prefix=pad_cot_prefix_lst, # CoT reuse + next_segment_llm_action=pad_llm_action_lst + ) + else: + last_game_segments[i].pad_over( + pad_obs_lst, pad_reward_lst, pad_action_lst, pad_root_values_lst, pad_child_visits_lst, + next_segment_raw_obs=pad_raw_obs_lst, next_segment_history_obs=pad_history_obs_lst, + next_segment_llm_prior_per_tok=pad_llm_prior_per_tok_lst, + next_segment_cot_prefix=pad_cot_prefix_lst, # CoT reuse + next_segment_llm_action=pad_llm_action_lst + ) + + last_game_segments[i].game_segment_to_array() + + # Add the completed game segment to the pool. + self.game_segment_pool.append((last_game_segments[i], last_game_priorities[i], done[i])) + + # Reset placeholders for the next collection cycle. + last_game_segments[i] = None + last_game_priorities[i] = None + + def collect( + self, + num_segments: Optional[int] = None, + train_iter: int = 0, + policy_kwargs: Optional[dict] = None, + collect_with_pure_policy: bool = False, + phase: Optional[str] = None + ) -> List[Any]: + """ + [PRIORZERO-MODIFIED] + Collect game segments with LLM-guided MCTS. + + Main changes from parent: + 1. Extract text observations from environment + 2. Pass LLM priors to policy forward pass + 3. Update history buffers after each step + + Args: + num_segments: Number of segments to collect + train_iter: Current training iteration + policy_kwargs: Additional kwargs for policy + collect_with_pure_policy: Whether to use pure policy without MCTS + + Returns: + return_data: List containing [game_segments, metadata] + """ + if num_segments is None: + if self._default_num_segments is None: + raise RuntimeError("Please specify num_segments for collection.") + else: + num_segments = self._default_num_segments + + assert num_segments == self._env_num, \ + f"num_segments({num_segments}) must equal env_num({self._env_num})" + + if policy_kwargs is None: + policy_kwargs = {} + + temperature = policy_kwargs.get('temperature', 1.0) + epsilon = policy_kwargs.get('epsilon', 0.0) + + collected_episode = 0 + collected_step = 0 + llm_prior_entropy = [[] for _ in range(self._env_num)] + env_nums = self._env_num + init_obs = self._env.ready_obs + + retry_waiting_time = 0.05 + while len(init_obs.keys()) != env_nums: + self._logger.info(f'[RANK {self._rank}] Waiting for all environments to reset. Ready: {list(init_obs.keys())}') + time.sleep(retry_waiting_time) + init_obs = self._env.ready_obs + + for env_id in range(env_nums): + if env_id in init_obs: + self.action_mask_dict[env_id] = to_ndarray(init_obs[env_id]['action_mask']) + self.to_play_dict[env_id] = to_ndarray(init_obs[env_id]['to_play']) + self.timestep_dict[env_id] = to_ndarray(init_obs[env_id].get('timestep', -1)) + + last_game_segments = [None for _ in range(env_nums)] + last_game_priorities = [None for _ in range(env_nums)] + game_segments = [ + GameSegment( + self._env.action_space, + game_segment_length=self.policy_config.game_segment_length, + config=self.policy_config, + task_id=self.task_id + ) for _ in range(env_nums) + ] + + observation_window_stack = [ + deque(maxlen=self.policy_config.model.frame_stack_num) + for _ in range(env_nums) + ] + for env_id in range(env_nums): + initial_frames = [ + to_ndarray(init_obs[env_id]['observation']) + for _ in range(self.policy_config.model.frame_stack_num) + ] + observation_window_stack[env_id].extend(initial_frames) + game_segments[env_id].reset(observation_window_stack[env_id], init_raw_obs=extract_raw_obs_text(init_obs[env_id]), + init_history_obs=list(self.history_buffers[env_id])) + + search_values_lst = [[] for _ in range(env_nums)] + pred_values_lst = [[] for _ in range(env_nums)] + + eps_steps_lst = np.zeros(env_nums) + visit_entropies_lst = np.zeros(env_nums) + llm_weight_lst = np.zeros(env_nums) + + if collect_with_pure_policy: + temp_visit_list = [0.0 for _ in range(self._env.action_space.n)] + + return_data = None + while True: + local_done = len(self.game_segment_pool) >= self._default_num_segments + + if local_done and return_data is None: + # First moment this rank reaches its target: log, snapshot + # return_data, clear pool. Do not break yet — TP partners on + # other DDP ranks may still be inside vllm.generate. + self._logger.info( + f'[RANK {self._rank}] ✓ Collected {len(self.game_segment_pool)} segments ' + f'(target: {self._default_num_segments})' + ) + return_data = [ + [self.game_segment_pool[i][0] for i in range(len(self.game_segment_pool))], + [ + { + 'priorities': self.game_segment_pool[i][1], + 'done': self.game_segment_pool[i][2], + 'unroll_plus_td_steps': self.unroll_plus_td_steps + } + for i in range(len(self.game_segment_pool)) + ] + ] + self.game_segment_pool.clear() + + if not self._should_continue_collect(local_done): + break + + if local_done: + # Drain mode: issue matched empty vllm iterations so TP partners + # don't deadlock at the next _sync_prompts_for_tp. + self.data_processor.drain_vllm_iter() + continue + + with self._timer: + obs = self._env.ready_obs + ready_env_id = set(obs.keys()) + + if len(ready_env_id) < self._env_num: + self._logger.debug(f'Only {len(ready_env_id)}/{self._env_num} envs ready') + + stack_obs_dict = { + env_id: game_segments[env_id].get_obs() + for env_id in ready_env_id + } + stack_obs_list = [stack_obs_dict[env_id] for env_id in sorted(list(ready_env_id))] + + action_mask = [self.action_mask_dict[env_id] for env_id in sorted(list(ready_env_id))] + to_play = [self.to_play_dict[env_id] for env_id in sorted(list(ready_env_id))] + timestep = [self.timestep_dict[env_id] for env_id in sorted(list(ready_env_id))] + + # Convert to tensors + stack_obs_array = to_ndarray(stack_obs_list) + stack_obs_tensor = prepare_observation( + stack_obs_array, + self.policy_config.model.model_type + ) + stack_obs_tensor = torch.from_numpy(stack_obs_tensor).to(self.policy_config.device) + + if collect_with_pure_policy: + continue + + # Extract text observations and valid actions + raw_obs_list = [] + histories_list = [] + valid_actions_list = [] + ready_env_ids = sorted(list(ready_env_id)) + for env_id in ready_env_ids: + raw_obs_text = extract_raw_obs_text(obs[env_id]) + raw_obs_list.append(raw_obs_text) + + history = list(self.history_buffers[env_id]) + histories_list.append(history) + + valid_actions = obs[env_id].get('valid_actions', []) + valid_actions_list.append(valid_actions) + with self.prof.block("collect_step_get_llm_prior", rank=self._rank): + # CoT reuse optimization: request CoT prefixes to store in game segments + llm_prior_per_seq, llm_prior_per_tok, cot_prefixes, _ = self.data_processor.get_llm_prior( + states=raw_obs_list, + valid_actions_list=valid_actions_list, # [PRIORZERO] Pass valid actions + histories=histories_list, + return_cot=True # Request CoT prefixes for reuse in training + ) + assert len(llm_prior_per_seq) == len(ready_env_id) == len(valid_actions_list) + for idx, llm_prior in enumerate(llm_prior_per_seq): + scaled_llm_prior = self.apply_temperature_scaling(llm_prior, return_logprobs=True) + llm_prior_per_seq[idx] = scaled_llm_prior + + llm_prior_per_seq_by_env = { + env_id: llm_prior_per_seq[idx] for idx, env_id in enumerate(ready_env_ids) + } + llm_prior_per_tok_by_env = { + env_id: llm_prior_per_tok[idx] for idx, env_id in enumerate(ready_env_ids) + } + cot_prefixes_by_env = { + env_id: cot_prefixes[idx] for idx, env_id in enumerate(ready_env_ids) + } + + policy_kwargs_forward = { + 'llm_prior_logprob': llm_prior_per_seq, + 'valid_actions_list': valid_actions_list, + "current_env_step": self._total_envstep_count, + "phase": phase, + "llm_collect_mode": self.llm_cfg.train_schedule['llm_collect_mode'] + } + + if self.task_id is not None: + policy_kwargs_forward['task_id'] = self.task_id + with self.prof.block("collect_step_forward", rank=self._rank): + policy_output = self._policy.forward(data=stack_obs_tensor, action_mask=action_mask, + temperature=temperature, to_play=to_play, epsilon=epsilon, + ready_env_id=sorted(list(ready_env_id)), timestep=timestep, + **policy_kwargs_forward) + + # Extract outputs + actions_with_env_id = {k: v['action'] for k, v in policy_output.items()} + value_dict_with_env_id = {k: v['searched_value'] for k, v in policy_output.items()} + pred_value_dict_with_env_id = {k: v['predicted_value'] for k, v in policy_output.items()} + + distributions_dict_with_env_id = { + k: v['visit_count_distributions'] for k, v in policy_output.items() + } + visit_entropy_dict_with_env_id = { + k: v['visit_count_distribution_entropy'] for k, v in policy_output.items() + } + llm_weight_dict_with_env_id = {k: v.get('llm_weight', 0.0) for k, v in policy_output.items()} + + actions: Dict[int, Any] = { + env_id: actions_with_env_id.pop(env_id) + for env_id in ready_env_id + } + with self.prof.block("collect_step", rank=self._rank): + timesteps = self._env.step(actions) + + interaction_duration = self._timer.value / len(timesteps) + + for env_id, episode_timestep in timesteps.items(): + with self._timer: + # Handle abnormal timesteps + if episode_timestep.info.get('abnormal', False): + self._env.reset({env_id: None}) + self._policy.reset([env_id]) + self._reset_stat(env_id) + self._logger.info(f'[RANK {self._rank}] Env {env_id} had abnormal step: {episode_timestep.info}') + continue + + obs_new, reward, done, info = ( + episode_timestep.obs, + episode_timestep.reward, + episode_timestep.done, + episode_timestep.info + ) + game_segments[env_id].store_search_stats( + distributions_dict_with_env_id[env_id], + value_dict_with_env_id[env_id]) + # =========================================================== + # [PRIORZERO-NEW] Update History Buffer + # =========================================================== + raw_obs_text = extract_raw_obs_text(obs[env_id]) + action = info['action_str'] + self.history_buffers[env_id].append((raw_obs_text, action, float(reward))) + + # Append transition to game segment (including CoT prefix for reuse optimization) + game_segments[env_id].append( + actions[env_id], + to_ndarray(obs_new['observation']), + reward, + self.action_mask_dict[env_id], + self.to_play_dict[env_id], + timestep=to_ndarray(self.timestep_dict[env_id]), + raw_obs_text=extract_raw_obs_text(obs_new), + history_obs=list(self.history_buffers[env_id]), + llm_prior_per_tok=llm_prior_per_tok_by_env[env_id], + cot_prefix=cot_prefixes_by_env[env_id], + llm_action=action + ) + + # Update state + self.action_mask_dict[env_id] = to_ndarray(obs_new['action_mask']) + self.to_play_dict[env_id] = to_ndarray(obs_new['to_play']) + self.timestep_dict[env_id] = to_ndarray(obs_new.get('timestep', -1)) + self.dones[env_id] = False if self.policy_config.ignore_done else done + + if not collect_with_pure_policy: + visit_entropies_lst[env_id] += visit_entropy_dict_with_env_id[env_id] + llm_weight_lst[env_id] += llm_weight_dict_with_env_id[env_id] + + eps_steps_lst[env_id] += 1 + + # Reset policy if needed (for UniZero) + if self._policy.get_attribute('cfg').type in ['unizero', 'sampled_unizero', 'priorzero']: + self._policy.reset( + env_id=env_id, + current_steps=eps_steps_lst[env_id], + reset_init_data=False + ) + + # Store values for priority calculation + if self.policy_config.use_priority: + pred_values_lst[env_id].append(pred_value_dict_with_env_id[env_id]) + search_values_lst[env_id].append(value_dict_with_env_id[env_id]) + + # Update observation window + observation_window_stack[env_id].append(to_ndarray(obs_new['observation'])) + + # =========================================================== + # Save Full Game Segment + # =========================================================== + if game_segments[env_id].is_full(): + if last_game_segments[env_id] is not None: + self.pad_and_save_last_trajectory(env_id, last_game_segments, last_game_priorities, + game_segments, self.dones) + + # Calculate priorities + priorities = self._compute_priorities(env_id, pred_values_lst, search_values_lst) + pred_values_lst[env_id], search_values_lst[env_id] = [], [] + + # Save segment + last_game_segments[env_id] = game_segments[env_id] + last_game_priorities[env_id] = priorities + + # Create new segment + game_segments[env_id] = GameSegment( + self._env.action_space, + game_segment_length=self.policy_config.game_segment_length, + config=self.policy_config, + task_id=self.task_id + ) + game_segments[env_id].reset(observation_window_stack[env_id], init_raw_obs=extract_raw_obs_text(obs_new), init_history_obs=list(self.history_buffers[env_id])) + + self._env_info[env_id]['step'] += 1 + if llm_prior_per_seq is not None and llm_prior_per_seq_by_env[env_id] is not None: + llm_prior_tensor = torch.tensor([logit for k, logit in llm_prior_per_seq_by_env[env_id].items()]) + llm_prior_prob = torch.softmax(llm_prior_tensor, dim=-1) + llm_prior_entropy[env_id].append(-torch.sum(llm_prior_prob * torch.log(llm_prior_prob + 1e-9), dim=-1)) + else: + llm_prior_entropy[env_id].append(0.0) + collected_step += 1 + + self._env_info[env_id]['time'] += self._timer.value + interaction_duration + + # ============================================================== + # Episode Done + # ============================================================== + if episode_timestep.done: + self._logger.info(f'[RANK {self._rank}] ======== Env {env_id} episode finished! ========') + # Logging + info_log = { + 'reward': episode_timestep.info['score'], + 'time': self._env_info[env_id]['time'], + 'step': self._env_info[env_id]['step'], + 'llm_prior_entropy': sum(llm_prior_entropy[env_id])/len(llm_prior_entropy[env_id])} + + self._logger.info( + f"[RANK {self._rank}] [Episode Complete] Env={env_id} | " + f"Reward={info_log['reward']:.2f} | " + f"Steps={info_log['step']} | " + f"Time={info_log['time']:.2f}s | " + f"LLM_Entropy={info_log['llm_prior_entropy']:.3f}" + ) + + if not collect_with_pure_policy: + info_log['visit_entropy'] = ( + visit_entropies_lst[env_id] / eps_steps_lst[env_id] + if eps_steps_lst[env_id] > 0 else 0 + ) + info_log['llm_weight'] = llm_weight_lst[env_id] / eps_steps_lst[env_id] if eps_steps_lst[env_id] > 0 else 0 + + + collected_episode += 1 + self._episode_info.append(info_log) + # Save remaining segments + if last_game_segments[env_id] is not None: + self.pad_and_save_last_trajectory( env_id, last_game_segments, last_game_priorities, game_segments, self.dones) + + priorities = self._compute_priorities( env_id, pred_values_lst, search_values_lst) + game_segments[env_id].game_segment_to_array() + if len(game_segments[env_id].reward_segment) > 0: + self.game_segment_pool.append(( + game_segments[env_id], + priorities, + self.dones[env_id] + )) + # Reset + pred_values_lst[env_id], search_values_lst[env_id] = [], [] + eps_steps_lst[env_id], visit_entropies_lst[env_id] = 0, 0 + llm_weight_lst[env_id] = 0 + + self._policy.reset([env_id], task_id=self.task_id) + self._reset_stat(env_id) + + # Clear history buffer for this environment + self.history_buffers[env_id].clear() + # Re-initialize game segment + init_obs = self._env.ready_obs + observation_window_stack[env_id] = deque( + [init_obs[env_id]['observation'] for _ in range(self.policy_config.model.frame_stack_num)], + maxlen=self.policy_config.model.frame_stack_num + ) + + game_segments[env_id] = GameSegment( + self._env.action_space, + game_segment_length=self.policy_config.game_segment_length, + config=self.policy_config, + task_id=self.task_id + ) + game_segments[env_id].reset(observation_window_stack[env_id], init_raw_obs=extract_raw_obs_text(init_obs[env_id]), init_history_obs=list(self.history_buffers[env_id])) + last_game_segments[env_id] = None + last_game_priorities[env_id] = None + + # ================================================================== + # Final Logging + # ================================================================== + collected_duration = sum([d['time'] for d in self._episode_info]) + + if self._world_size > 1: + # Before allreduce + local_step, local_episode = collected_step, collected_episode + collected_step = allreduce_data(collected_step, 'sum') + collected_episode = allreduce_data(collected_episode, 'sum') + collected_duration = float(collected_duration) + collected_duration = allreduce_data(collected_duration, 'sum') + # After allreduce + self._logger.info( + f"[Rank {self._rank} Aggregation] " + f"Local: steps={local_step}, episodes={local_episode} | " + f"Global: steps={collected_step}, episodes={collected_episode}" + ) + + self._total_envstep_count += collected_step + self._total_episode_count += collected_episode + self._total_duration += collected_duration + + self._output_log(train_iter) + + return return_data + + def _output_log(self, train_iter: int) -> None: + """ + [INHERITED] + Log collection statistics (inherited from parent). + """ + if self._rank != 0: + return + + if (train_iter - self._last_train_iter) >= self._collect_print_freq and len(self._episode_info) > 0: + self._last_train_iter = train_iter + episode_count = len(self._episode_info) + envstep_count = sum([d['step'] for d in self._episode_info]) + duration = sum([d['time'] for d in self._episode_info]) + episode_reward = [d['reward'] for d in self._episode_info] + episode_llm_prior_entropy = [d['llm_prior_entropy'] for d in self._episode_info] + + info = { + 'episode_count': episode_count, + 'envstep_count': envstep_count, + 'avg_envstep_per_episode': envstep_count / episode_count, + 'avg_envstep_per_sec': envstep_count / duration if duration > 0 else 0, + 'avg_episode_per_sec': episode_count / duration if duration > 0 else 0, + 'collect_time': duration, + 'reward_mean': np.mean(episode_reward), + 'reward_std': np.std(episode_reward), + 'reward_max': np.max(episode_reward), + 'reward_min': np.min(episode_reward), + 'total_envstep_count': self._total_envstep_count, + 'total_episode_count': self._total_episode_count, + 'total_duration': self._total_duration, + 'llm_prior_entropy_mean': np.mean(episode_llm_prior_entropy), + 'llm_prior_entropy_max': np.max(episode_llm_prior_entropy), + 'llm_prior_entropy_min': np.min(episode_llm_prior_entropy) + } + + if not self.collect_with_pure_policy: + visit_entropy = [d['visit_entropy'] for d in self._episode_info] + info['visit_entropy_mean'] = np.mean(visit_entropy) + llm_weight = [d['llm_weight'] for d in self._episode_info] + info['llm_weight_mean'] = np.mean(llm_weight) + if self.policy_config.gumbel_algo: + completed_value = [d['completed_value'] for d in self._episode_info] + info['completed_value_mean'] = np.mean(completed_value) + + self._episode_info.clear() + + self._logger.info( + f"\n{'='*80}\n" + f"[RANK {self._rank}][Collector Summary] Train Iter: {train_iter}\n" + f"{'-'*80}\n" + f"Episodes: {info['episode_count']} (Total: {info['total_episode_count']})\n" + f"Steps: {info['envstep_count']} (Total: {info['total_envstep_count']})\n" + f"Avg Steps/Ep: {info['avg_envstep_per_episode']:.1f}\n" + f"Throughput: {info['avg_envstep_per_sec']:.2f} steps/s, {info['avg_episode_per_sec']:.3f} eps/s\n" + f"Duration: {info['collect_time']:.2f}s (Total: {info['total_duration']:.2f}s)\n" + f"{'-'*80}\n" + f"Reward: mean={info['reward_mean']:.2f}, std={info['reward_std']:.2f}, " + f"min={info['reward_min']:.2f}, max={info['reward_max']:.2f}\n" + f"LLM Entropy: mean={info['llm_prior_entropy_mean']:.3f}, " + f"min={info['llm_prior_entropy_min']:.3f}, max={info['llm_prior_entropy_max']:.3f}\n" + + (f"Visit Entropy: {info.get('visit_entropy_mean', 0):.3f}\n" if not self.collect_with_pure_policy else "") + + (f"Completed Val: {info.get('completed_value_mean', 0):.3f}\n" if self.policy_config.gumbel_algo else "") + + f"{'='*80}" + ) + + # Log to console + self._logger.info("Collector Training Summary:\n{}".format('\n'.join([f' {k}: {v}' for k, v in info.items()]))) + + # Log to TensorBoard and WandB + for k, v in info.items(): + if self.task_id is None: + tb_prefix_iter = f'{self._instance_name}_iter/' + tb_prefix_step = f'{self._instance_name}_step/' + else: + tb_prefix_iter = f'{self._instance_name}_iter_task{self.task_id}/' + tb_prefix_step = f'{self._instance_name}_step_task{self.task_id}/' + + self._tb_logger.add_scalar(tb_prefix_iter + k, v, train_iter) + self._tb_logger.add_scalar(tb_prefix_step + k, v, self._total_envstep_count) + + def apply_temperature_scaling(self, logprobs_dict: dict, return_logprobs: bool = True) -> dict: + """ + 对 Logprobs 字典进行温度缩放,控制分布的平缓程度。 + """ + T = self.llm_prior_temperature + if T <= 1e-8: + max_key = max(logprobs_dict, key=logprobs_dict.get) + return {k: (0.0 if k != max_key else 1.0) for k in logprobs_dict} + + scaled_logits = {k: v / T for k, v in logprobs_dict.items()} + + max_val = max(scaled_logits.values()) + sum_exp = sum(math.exp(v - max_val) for v in scaled_logits.values()) + log_sum_exp = math.log(sum_exp) + max_val + + result = {} + for k, v in scaled_logits.items(): + normalized_logprob = v - log_sum_exp + + if return_logprobs: + result[k] = normalized_logprob + else: + result[k] = math.exp(normalized_logprob) + + return result diff --git a/zoo/jericho/priorzero/src/priorzero_config.py b/zoo/jericho/priorzero/src/priorzero_config.py new file mode 100644 index 000000000..6841a326a --- /dev/null +++ b/zoo/jericho/priorzero/src/priorzero_config.py @@ -0,0 +1,491 @@ +import os +from typing import Dict, Tuple, Optional, Any +from easydict import EasyDict +import torch.distributed as dist +from dataclasses import dataclass, field + +# ============================================================================ +# Model Configuration Presets +# ============================================================================ +MODEL_CONFIGS = { + "qwen2.5-0.5b": { + "model_name_or_path": "/mnt/shared-storage-user/puyuan/xiongjyu/models/Qwen2.5-0.5B-Instruct", + "vllm_tensor_parallel_size": 1, + "gpu_memory_utilization": 0.2, + "description": "Qwen2.5-0.5B-Instruct (smallest, fastest)", + }, + "qwen2.5-1.5b": { + "model_name_or_path": "/mnt/shared-storage-user/puyuan/xiongjyu/models/Qwen2.5-1.5B-Instruct", + "vllm_tensor_parallel_size": 1, + "gpu_memory_utilization": 0.2, + "description": "Qwen2.5-1.5B-Instruct (balanced performance)", + }, + "qwen2.5-3b": { + # "model_name_or_path": "/mnt/afs/niuyazhe/workspace/xiongjyu/models/Qwen2.5-3B-Instruct", + "model_name_or_path": "/mnt/shared-storage-user/puyuan/xiongjyu/models/Qwen2.5-3B-Instruct", + "vllm_tensor_parallel_size": 1, + "gpu_memory_utilization": 0.25, + "description": "Qwen2.5-3B-Instruct (better quality)", + }, + "qwen2.5-7b": { + # "model_name_or_path": "/mnt/shared-storage-user/puyuan/model/Qwen2.5-7B-Instruct", + # "vllm_tensor_parallel_size": 2, + "model_name_or_path": "/mnt/shared-storage-user/puyuan/xiongjyu/models/Qwen2.5-7B-Instruct", + "vllm_tensor_parallel_size": 1, + + "gpu_memory_utilization": 0.35, + "description": "Qwen2.5-7B-Instruct (high quality, needs 2+ GPUs)", + }, + "qwen2.5-14b": { + "model_name_or_path": "/mnt/shared-storage-user/puyuan/model/Qwen2.5-14B-Instruct", + "vllm_tensor_parallel_size": 4, + "gpu_memory_utilization": 0.5, + "description": "Qwen2.5-14B-Instruct (best quality, needs 4+ GPUs)", + }, +} + +def get_available_models(): + """Get list of available model configurations""" + return list(MODEL_CONFIGS.keys()) + +def get_model_config(model_key: str) -> Dict: + """Get model configuration by key""" + if model_key not in MODEL_CONFIGS: + available = ", ".join(get_available_models()) + raise ValueError( + f"Unknown model key: {model_key}\n" + f"Available models: {available}" + ) + return MODEL_CONFIGS[model_key] + +def print_available_models(): + """Print all available model configurations""" + print("\n" + "="*80) + print("Available Model Configurations:") + print("="*80) + for key, config in MODEL_CONFIGS.items(): + print(f"\n {key}:") + print(f" Path: {config['model_name_or_path']}") + print(f" Tensor Parallel Size: {config['vllm_tensor_parallel_size']}") + print(f" GPU Memory Utilization: {config['gpu_memory_utilization']}") + print(f" Description: {config['description']}") + print("="*80 + "\n") + +@dataclass +class PriorZeroLLMConfig: + model_name_or_path: str = "Qwen2.5-3B-Instruct" + enable_rft: bool = True + enable_world_model: bool = True + train_mode_dict: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "mode": "full", # "full" or "lora" + "lora_r": 16, + "lora_alpha": 32, + "lora_dropout": 0.05, + "lora_bias": "none", # "none" / "all" / "lora_only" + "lora_target_modules": ( + "q_proj", + "k_proj", + "v_proj", + "o_proj", + "gate_proj", + "up_proj", + "down_proj", + ), + })) + + train_schedule: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "alternate": True, # False 两者都训练(默认配置);True: 严格交替训练:phase=wm 时仅训练 wm;phase=llm 时仅训练 llm + "wm_update_iters": 2e3, # alternate=True. wm 的 train_iter + "llm_update_iters": 2e2, # alternate=True. llm 的 train_iter + "start_phase": "wm", # alternate=True. 从哪个阶段开始: "wm" 或 "llm" + "llm_collect_mode": "no_collect" # wm_collect意味着llm训练过程收集数据使用 wm; wm_llm_collect意味着 llm 训练过程收集数据使用 llm 和 wm; no_collect 意味着 llm 训练过程不收集数据,直接使用 replay buffer 中的数据 + })) + + llm_prior_temperature: float = 2.0 # LLM prior 分布的温度参数 + mcts_root_logits_dict: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "mode": "llm_plus_wm_logits", # collect/eval阶段保持一致。"llm_logits"是仅用llm prior的logits; "wm_logits"是仅用 world_model 的policy给出的logits; "llm_plus_wm_logits"是两者的加权求和。 + "plus_method": "fixed", # 当 plus_method = "fixed" 时,使用固定权重;否则使用自适应权重"adaptive" + "wm_weight": 0.5, # 当 plus_method = "fixed" 时,WM logits 的权重;LLMPrior 的权重 = 1 - WM_weight + "llm_max_weight": 0.7, # 当 plus_method = "adaptive" 时,LLM 的最大权重;WM 的最小权重 = 1 - llm_max_weight + "llm_min_weight": 0.3, + "max_envsteps": 1e5, # 当 plus_method = "adaptive" 时,随着环境交互步数增加,逐渐降低 llm prior 的权重 + })) + eval_dict: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "world_model": True, # 评估模式1:完全与 unizero 的 eval 一致;mcts 的根节点仅使用 WM 的logits + "world_model_llm_prior": True, # 评估模式2:基于 unizero 的 eval 过程, 但是 mcts 的根节点需要利用 llm 的先验;具体怎么利用取决于mcts_root_logits_dict.mode 参数 + "llm_prior": True, # 评估模式3:仅使用 llm prior 进行 eval, 不需要 wm 进行评估 + "wm_eval_freq": 500, + "llm_eval_freq": 50, + # env-step-based eval frequency (preferred over iter-based when > 0) + "wm_eval_freq_envsteps": 0, # 0 = disabled, falls back to wm_eval_freq + "llm_eval_freq_envsteps": 0, # 0 = disabled, falls back to llm_eval_freq + # Whether to save CoT/prompt/LLM-prior details in trajectory JSON files + "save_llm_cot": True, + })) + + attn_implementation: str = "flash_attention_2" + history_length: int = 25 + use_cot: bool = False + cot_weight: float = 0.1 # 控制 cot前缀token的权重,由于重点是action:,所以前缀的token权重调低 + + user_prompt_dict: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "history_with_reward": True, # 是否在 prompt 中加入历史交互的 reward 信息 + "observation_with_valid_actions": False, # 是否在 prompt 中加入当前 observation 中可执行的 action 信息 + })) + + prompt_max_len: int = 8192 + generate_max_len: int = 512 + bf16: bool = True + + # vLLM engines + enable_vllm: bool = True + enable_prefix_caching: bool = False + use_cuda_ipc: bool = False + enable_vllm_is_correction: bool = False + vllm_is_truncated_threshold: Tuple[float, float] = (0.5, 5.0) + use_mispo: bool = False + mispo_token_truncated_threshold: Tuple[float, float] = (0.5, 2.0) + mispo_traj_truncated_threshold: Tuple[float, float] = (0.8, 1.2) + + vllm_sync_backend: str = "nccl" # vLLM 同步参数使用的后端 + vllm_tensor_parallel_size: int = 1 # 每个vllm engine使用几张GPU张量并行 (Fixed: 1.5B model should use 1 GPU) + + gpu_memory_utilization: float = 0.3 + vllm_enable_sleep: bool = True # 是否可以休眠 + temperature: float = 1.0 + top_p: float = 0.95 + seed: int = 0 + reduction: str = "mean" + + # 训练相关参数 + deepspeed_enable_sleep: bool = True + + zero_stage: int = 2 + gradient_checkpointing: bool = False + gradient_checkpointing_use_reentrant: bool = False + max_norm: float = 1.0 # Gradient clipping + ds_tensor_parallel_size: int = 1 + + # 需要注意的是,buffer中取一条经验是 10个样本,因为包含10次交互; num_unroll_steps = 10 + train_batch_size: int = 128 # 总的train_size, 结果= micro_batch_size * GPUS * gradient_accumulation_steps + micro_train_batch_size: int = 4 # 一次micro_train_batch_size 用来计算梯度;只有一次 train_batch_size 才会更新参数 + max_rollout_staleness: int = 1 # off 次数,用来训练的数据和当前策略之间允许的最大差距 + + learning_rate: float = 1e-6 + adam_betas: Tuple[float, float] = (0.9, 0.95) + weight_decay: float = 0.01 + lr_scheduler: str = "cosine_with_min_lr" + lr_warmup_ratio: float = 0.03 + max_steps: int = int(1e4) + policy_loss_type: str = "ppo" # 'ppo' / 'gspo' + reward_func: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "format_reward": True, + "format_param": EasyDict( + {"format_weight": 0.5, } # fmt_reward 的权重,应该在 [0, 1) 之间,因为advantage的权重是 1 - format_weight + ), + })) + # advantage = target_value - pred_value + # advantage_global_batch_norm:意味着 llm训练阶段,所有训练数据的 advantage + # advantage_batch_norm:意味着 llm 训练过程,train_batch_size之前取advantage + advantage_type: str = "advantage_global_batch_norm" # "advantage", "target_reward", "advantage_batch_norm", "advantage_running_norm" "advantage_global_batch_norm" + eps_clip_low_high: Tuple[float, float] = (0.2, 0.2) + rft_kl_coef: float = 0.01 + entropy_loss_coef: float = 0.0 + kl_estimator: str = "k3" + kl_early_stop_threshold: float = 0.0 # 0 means disabled; when ref_kl exceeds this, skip remaining gradient updates in the epoch + + llm_save_freq: int = 1000 # 每多少步保存一次 llm 模型,一步代表一次参数更新而不是梯度累积 + save_path: str = "" # 该参数将被 exp_name 目录覆盖 + + value_norm_cfg: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + 'enable_stability_optimizer': True, + 'value_norm_init_momentum': 0.9, # Fast adaptation in early training + 'value_norm_final_momentum': 0.99, # Slow, stable updates in later training + 'value_norm_warmup_steps': 100, # Steps to transition from init to final momentum + 'value_norm_clip_percentile': 0.95, # Clip outliers beyond this percentile + 'value_norm_clip_method': "soft", + "value_norm_history_size": 1000, + })) + + +def get_priorzero_config( + env_id: str = 'detective.z5', + seed: int = 0, + exp_name: str = None, + use_cot: bool = False, + model_key: Optional[str] = "qwen2.5-3b", + multi_gpu: bool = False, +) -> Tuple[EasyDict, EasyDict]: + """ + Generate complete PriorZero configuration with automatic model configuration. + + Args: + env_id: Jericho game ID + seed: Random seed + exp_name: Experiment name (auto-generated if None) + use_cot: Whether to use Chain-of-Thought reasoning + model_key: Model configuration key (e.g., 'qwen2.5-0.5b', 'qwen2.5-1.5b', 'qwen2.5-7b') + If None, uses default 'qwen2.5-1.5b' + + Returns: + main_config: Main configuration dictionary + create_config: Creation configuration for DI-engine components + llm_config: LLM configuration with auto-configured model parameters + """ + env_configurations = { + 'detective.z5': (12, 100), + 'omniquest.z5': (25, 100), + 'acorncourt.z5': (45, 50), + 'zork1.z5': (55, 500), + } + action_space_size, max_steps = env_configurations.get(env_id, (20, 100)) + wm_encoder_option = 'legacy' + # wm_model_name = 'BAAI/bge-base-en-v1.5' + # wm_model_name = '/mnt/afs/niuyazhe/workspace/xiongjyu/models/bge-base-en-v1.5' + wm_model_name = '/mnt/shared-storage-user/puyuan/xiongjyu/models/bge-base-en-v1.5' + + collector_env_num = 1 + evaluator_env_num = 2 + n_episode = collector_env_num + + num_unroll_steps = 10 + infer_context_length = 4 + game_segment_length = 50 + num_layers = 2 + embed_dim = 768 + replay_ratio = 0.1 + batch_size = 64 + collect_num_simulations=25 + eval_num_simulations=25 + replay_buffer_size = int(3e5) + + env_config = dict( + stop_value=int(1e6), + max_steps=max_steps, + observation_shape=512, + env_id=env_id, + # game_path=f"/mnt/afs/wanzunian/niuyazhe/xiongjyu/jericho/LightZero/zoo/jericho/envs/z-machine-games-master/jericho-game-suite/{env_id}", + # game_path=f"/mnt/afs/niuyazhe/workspace/xiongjyu/LightZero/zoo/jericho/envs/z-machine-games-master/jericho-game-suite/{env_id}", + game_path=f"/mnt/shared-storage-user/puyuan/code/LightZero/zoo/jericho/envs/z-machine-games-master/jericho-game-suite/{env_id}", + for_unizero=True, + tokenizer_path=wm_model_name, + max_action_num=action_space_size, + max_seq_len=512, + collector_env_num=collector_env_num, + evaluator_env_num=evaluator_env_num, + n_evaluator_episode=evaluator_env_num, + manager=dict( + shared_memory=False, + ), + use_cache=True, + cache_size=100000, + get_valid_actions_timeout=40 + ) + policy_config = dict( + type='priorzero', + multi_gpu=multi_gpu, + use_wandb=False, + learn=dict( + learner=dict( + hook=dict( + save_ckpt_after_iter=1000000, + ), + ), + ), + model=dict( + reward_support_range=(-300., 301., 1.), + value_support_range=(-300., 301., 1.), + + observation_shape=512, + action_space_size=action_space_size, + encoder_option=wm_encoder_option, + encoder_url=wm_model_name, + model_type="mlp", + continuous_action_space=False, + norm_type="LN", + world_model_cfg=dict( + support_size=601, + + norm_type="LN", + final_norm_option_in_head="LayerNorm", + final_norm_option_in_encoder="LayerNorm", + predict_latent_loss_type='mse', + policy_entropy_weight=5e-2, + continuous_action_space=False, + max_blocks=num_unroll_steps, + max_tokens=2 * num_unroll_steps, + context_length=2 * infer_context_length, + device="cuda", + action_space_size=action_space_size, + num_layers=num_layers, + num_heads=24, + embed_dim=embed_dim, + obs_type="text", + env_num=max(collector_env_num, evaluator_env_num), + decode_loss_mode=None, + latent_recon_loss_weight=0, + + task_embed_option=None, + moe_in_transformer=False, + multiplication_moe_in_transformer=False, + game_segment_length=game_segment_length, + ) + ), + update_per_collect=None, + num_segments=collector_env_num, + action_type="varied_action_space", + model_path=None, + num_unroll_steps=num_unroll_steps, + reanalyze_ratio=0, + replay_ratio=replay_ratio, + batch_size=batch_size, + learning_rate=3e-4, + weight_decay=1e-4, + cos_lr_scheduler=False, + fixed_temperature_value=0.25, + manual_temperature_decay=False, + n_episode=n_episode, + train_start_after_envsteps=0, + replay_buffer_size=replay_buffer_size, + eval_freq=int(3e4), + collector_env_num=collector_env_num, + evaluator_env_num=evaluator_env_num, + buffer_reanalyze_freq=1 / 1000000, + reanalyze_batch_size=160, + reanalyze_partition=0.75, + device='cuda', + + collect_num_simulations=collect_num_simulations, + eval_num_simulations=eval_num_simulations, + game_segment_length=game_segment_length, + off_policy_degree=0, + enable_async_eval=False, + + optim_type='AdamW', + grad_clip_value=10.0, + value_loss_weight=0.25, + policy_loss_weight=1.0, + reward_loss_weight=1.0, + + use_adaptive_entropy_weight=False, + adaptive_entropy_alpha_lr=1e-4, + use_encoder_clip_annealing=False, + encoder_clip_anneal_type='cosine', + encoder_clip_start_value=30.0, + encoder_clip_end_value=10.0, + encoder_clip_anneal_steps=100000, + use_priority=False, # Prioritized experience replay + priority_prob_alpha=0.6, + priority_prob_beta=0.4, + ) + + llm_config = PriorZeroLLMConfig(use_cot=use_cot) # 需要修改 llm 相关的参数,修改以上类即可 + + # Apply model configuration + model_config = get_model_config(model_key) + llm_config.model_name_or_path = model_config["model_name_or_path"] + llm_config.vllm_tensor_parallel_size = model_config["vllm_tensor_parallel_size"] + llm_config.gpu_memory_utilization = model_config["gpu_memory_utilization"] + + if exp_name is None: + env_name = env_id.replace(".z5", "") + if llm_config.enable_rft: + exp_name = ( + f"data_priorzero/llm_rft/priorzero_{env_name}_{model_key}_train_{llm_config.train_mode_dict.mode}/" + f"useCot_{llm_config.use_cot}_alternate_{llm_config.train_schedule.alternate}/" + f"mcts_{llm_config.mcts_root_logits_dict.mode}_staleness_{llm_config.max_rollout_staleness}_tbs_{llm_config.train_batch_size}_use_mispo_{llm_config.use_mispo}" + ) + else: + exp_name = ( + f"data_priorzero/llm_frozen/priorzero_{env_name}_{model_key}_" + f"train_{llm_config.train_mode_dict.mode}" + f"useCot_{llm_config.use_cot}_seed{seed}" + ) + + priorzero_config = dict( + env=env_config, + policy=policy_config, + exp_name=exp_name, + seed=seed + ) + create_config = dict( + env=dict( + type="jericho", + import_names=["zoo.jericho.envs.jericho_env"], + ), + env_manager=dict( + type="base" + ), + policy=dict( + type="priorzero", + import_names=["zoo.jericho.priorzero.src.priorzero_policy"], + ), + collector=dict( + type="priorzero_segment", + import_names=["zoo.jericho.priorzero.src.priorzero_collector"], + ), + evaluator=dict( + type="priorzero", + import_names=["zoo.jericho.priorzero.src.priorzero_evaluator"], + ), + replay_buffer=dict( + type='game_buffer_muzero', + import_names=['lzero.mcts.buffer.game_buffer_muzero'], + ), + ) + main_config = EasyDict(priorzero_config) + create_config = EasyDict(create_config) + + print(f"[Config] Model configuration applied:") + print(f" - Model: {model_key}") + print(f" - Path: {llm_config.model_name_or_path}") + print(f" - Train Mode: {llm_config.train_mode_dict.mode}") + print(f" - Tensor Parallel Size: {llm_config.vllm_tensor_parallel_size}") + print(f" - GPU Memory Utilization: {llm_config.gpu_memory_utilization}") + if llm_config.train_mode_dict.mode == "lora": + print( + f" - LoRA r/alpha/dropout: " + f"{llm_config.train_mode_dict.lora_r}/" + f"{llm_config.train_mode_dict.lora_alpha}/" + f"{llm_config.train_mode_dict.lora_dropout}" + ) + print(f" - LoRA target modules: {', '.join(llm_config.train_mode_dict.lora_target_modules)}") + + return main_config, create_config, llm_config + + +def get_priorzero_debug_config( + env_id: str = 'detective.z5', + seed: int = 0, + exp_name: str = None, + use_cot: bool = False, + model_key: Optional[str] = "qwen2.5-3b", +) -> EasyDict: + + main_config, create_config, llm_config = get_priorzero_config( + env_id=env_id, seed=seed, exp_name=exp_name, use_cot=use_cot, model_key=model_key + ) + max_steps = 20 + + batch_size = 8 + collect_num_simulations=2 + eval_num_simulations=2 + num_layers=1 + game_segment_length = 50 + + llm_config.train_batch_size = 8 # 总的train_size, 结果= micro_batch_size * GPUS * gradient_accumulation_steps + llm_config.micro_train_batch_size = 4 + llm_config.train_schedule.wm_update_iters=2 + llm_config.train_schedule.llm_update_iters=1 + + create_config.max_steps = max_steps + + main_config.policy.model.world_model_cfg.num_layers = num_layers + main_config.policy.model.world_model_cfg.game_segment_length = game_segment_length + main_config.policy.batch_size = batch_size + main_config.policy.collect_num_simulations = collect_num_simulations + main_config.policy.eval_num_simulations = eval_num_simulations + main_config.policy.update_per_collect = 2 + main_config.policy.game_segment_length = game_segment_length + + return main_config, create_config, llm_config diff --git a/zoo/jericho/priorzero/src/priorzero_datafactory.py b/zoo/jericho/priorzero/src/priorzero_datafactory.py new file mode 100644 index 000000000..abeaf8c1f --- /dev/null +++ b/zoo/jericho/priorzero/src/priorzero_datafactory.py @@ -0,0 +1,913 @@ +from __future__ import annotations +from dataclasses import dataclass +from typing import Any, Dict, List, Optional, Tuple + +import re +import torch +import torch.distributed as dist +from vllm import SamplingParams +from ding.utils import build_logger +import random +import numpy as np +import math +import logging + +_log_train = logging.getLogger("priorzero.train") + +_FMT_RE = re.compile( + r'^\s*Reasoning:\s*(?P[\s\S]*?)\nAction:\s*(?P[^\n\r]+)\s*$', + flags=re.IGNORECASE +) +def _format_reward(text: str) -> int: + """ + Return 1 if the output strictly matches: + Reasoning: + Action: + Otherwise 0. + """ + if not isinstance(text, str): + return 0 + + t = text.replace("\r\n", "\n").replace("\r", "\n").strip() + + m = _FMT_RE.match(t) + if m is None: + return 0 + + if len(re.findall(r'Reasoning:', t, flags=re.IGNORECASE)) != 1: + return 0 + if len(re.findall(r'Action:', t, flags=re.IGNORECASE)) != 1: + return 0 + + # Action 必须非空(regex 已经用 + 保证非空,这里再保险) + if m.group("action").strip() == "": + return 0 + + return 1 + + + +def unique_dicts_hash(lst): + import hashlib + import pickle + seen = set() + res = [] + for d in lst: + b = pickle.dumps(d) + h = hashlib.md5(b).hexdigest() + + if h not in seen: + seen.add(h) + res.append(d) + return res + +class DataProcessor: + """ + - build_llm_prompt / build_chat_context + - priorzero_batch -> samples + - (use_cot) 批量生成 prefix_cot + - vLLM 计算 action prior score(prompt_logprobs) + - samples -> Dataset/Dataloader(collate_fn 做 pack) + """ + + def __init__(self, rank, world_size, vllm_engine, strategy, model_path, exp_name=None, instance_name="vllm_output"): + self.vllm_engine = vllm_engine + self.strategy = strategy + self.args = getattr(strategy, "args", None) + + from transformers import AutoTokenizer + self.tokenizer = AutoTokenizer.from_pretrained( + model_path, trust_remote_code=True, padding_side="left" + ) + if self.tokenizer.pad_token is None: + self.tokenizer.pad_token = self.tokenizer.eos_token + + self.use_cot = self.args.use_cot + self.prompt_max_len = self.args.prompt_max_len + self.generate_max_len = self.args.generate_max_len + self.temperature = self.args.temperature + self.top_p = self.args.top_p + self.vllm_enable_sleep = self.args.vllm_enable_sleep + self.reduction = self.args.reduction + self.rank = rank + self.world_size = world_size + self.output_step = 0 + self.llm_prior_with_cot = False + + from collections import deque + self.episode_output = [] + + # Running statistics for advantage normalization + self.value_running_mean = 0.0 + self.value_running_std = 1.0 + self.value_count = 0 + self.running_momentum = 0.99 # EMA momentum for running statistics + + self.global_batch_advantages = [] + + if self.rank == 0: + self._logger, _ = build_logger( + path=f'./{exp_name}/log/{instance_name}', name=instance_name, need_tb=False + ) + + if self.args.value_norm_cfg.enable_stability_optimizer: + from models.stability_optimizer import AdaptiveValueNormalizer + self.value_normalizer = AdaptiveValueNormalizer( + init_momentum=self.args.value_norm_cfg.value_norm_init_momentum, + final_momentum=self.args.value_norm_cfg.value_norm_final_momentum, + warmup_steps=self.args.value_norm_cfg.value_norm_warmup_steps, + clip_method=self.args.value_norm_cfg.value_norm_clip_method, + clip_percentile=self.args.value_norm_cfg.value_norm_clip_percentile, + min_std=1e-6, + history_size=self.args.value_norm_cfg.value_norm_history_size, + ) + else: + self.value_normalizer = None + + def get_system_prompt(self): + """ + 系统提示词:纯文本指令,定义角色、目标和严格的输出协议。 + """ + parts = [ + "You are an expert player in a text-based adventure game. Your goal is to maximize the score by choosing the optimal next action.", + "Please analyze the game history and current observation to decide the single best next action.", + "OUTPUT FORMAT:", + ] + + if self.use_cot: + parts.append( + "You MUST produce exactly TWO parts in the following order:\n" + "1. Reasoning: Analyze the current situation, available actions, constraints, and uncertainties. Do NOT reveal the final choice here.\n" + "2. Action: The final chosen action.\n" + "Strict Format Example:\n" + "Reasoning: \n" + "Action: " + ) + else: + parts.append( + "Output exactly one line starting with 'Action:'.\n" + "Example:\n" + "Action: " + ) + return "\n".join(parts) + + def get_user_prompt( + self, + history: Optional[List[Tuple[str, str, float]]] = None, + current_obs: Optional[str] = None, + valid_actions: Optional[List[str]] = None + ) -> str: + """ + 用户提示词:注入历史和当前状态,并触发输出。 + """ + prompt_parts = [] + user_prompt_dict = self.args.user_prompt_dict + if history and len(history) > 0: + prompt_parts.append("=== GAME HISTORY ===") + for i, (obs, action, reward) in enumerate(history, start=1): + prompt_parts.append(f"Step {i}:") + prompt_parts.append(f"Observation: {obs.strip()}") + prompt_parts.append(f"Action: {action.strip()}") + if user_prompt_dict.history_with_reward: + prompt_parts.append(f"Reward: {reward}") + prompt_parts.append("") # 空行分隔 + + prompt_parts.append("=== CURRENT OBSERVATION ===") + prompt_parts.append(current_obs.strip()) + if user_prompt_dict.observation_with_valid_actions: + if valid_actions and len(valid_actions) > 0: + actions_str = ", ".join([f"'{act}'" for act in valid_actions]) + prompt_parts.append(f"\n[Valid Actions]\nYou can choose from the following actions: {actions_str}") + + prompt_parts.append("\n=== INSTRUCTION ===") + if self.use_cot: + prompt_parts.append( + "Please analyze the situation and provide your response in the following format:\n" + "Reasoning: \n" + "Action: " + ) + else: + prompt_parts.append( + "Decide on the best next move and output it in the following format:\n" + "Action: " + ) + return "\n".join(prompt_parts) + + def build_chat_context(self, user_prompt: str) -> str: + return self.tokenizer.apply_chat_template( + [ + {"role": "system", "content": self.get_system_prompt()}, + {"role": "user", "content": user_prompt} + ], + tokenize=False, + add_generation_prompt=True, + ) + + def build_llm_samples(self, + raw_obs_list: List[List[str]], + history_obs_list: List[List[List[Tuple[str, str, float]]]], + llm_prior_per_tok_list: Optional[List[List[Any]]] = None, + pred_values: Optional[torch.Tensor] = None, # [B, T-1] + target_values: Optional[torch.Tensor] = None, # [B, T-1] + cot_prefix_list: Optional[List[List[str]]] = None, # CoT reuse optimization + llm_action_list: Optional[List[List[str]]] = None, + ) -> List[Dict[str, Any]]: + """ + Build training samples from collected data. + + Args: + raw_obs_list: Raw observations + history_obs_list: History observations + llm_prior_per_tok_list: LLM prior per token from collect phase + target_values: Target values for advantage calculation + cot_prefix_list: CoT prefixes from collect phase (CoT reuse optimization) + + Returns: + List of sample dictionaries + """ + samples: List[Dict[str, Any]] = [] + B = len(raw_obs_list) + if B == 0: + return samples + T = len(raw_obs_list[0]) + + for b in range(B): + for t in range(T - 1): + current_obs = raw_obs_list[b][t] + current_hist = history_obs_list[b][t] + + true_action = llm_action_list[b][t+1] + rollout_logprob = llm_prior_per_tok_list[b][t+1]['rollout_action_logprob'][true_action] + full_ids = llm_prior_per_tok_list[b][t+1]['full_ids'][true_action] + label_ids = llm_prior_per_tok_list[b][t+1]['label_ids'][true_action] + valid_actions = list(llm_prior_per_tok_list[b][t+1]['rollout_action_logprob'].keys()) + if 'go' in valid_actions: + valid_actions.remove('go') + + instruction = self.get_user_prompt( + history=current_hist, + current_obs=current_obs, + valid_actions=valid_actions + ) + prompt = self.build_chat_context(instruction) + + if len(label_ids) == 0: + continue + target_value = None + if target_values is not None: + target_value = float(target_values[b][t].item()) + + pred_value = None + if pred_values is not None: + pred_value = float(pred_values[b][t].item()) + + # CoT reuse optimization: get CoT prefix from stored data + prefix_cot = None + if self.use_cot and cot_prefix_list is not None: + prefix_cot = cot_prefix_list[b][t+1] + + samples.append( + { + "instruction": instruction, + "prompt": prompt, + "target": true_action, + "pred_value": pred_value, + "target_value": target_value, + "rollout_logprob": rollout_logprob, # Reinforce++ ratio 需要 + "prefix_cot": prefix_cot, # CoT reuse optimization + "full_ids": full_ids, + "label_ids": label_ids, + } + ) + return samples + + def make_llm_train_samples(self, priorzero_batch, ddp: bool = False, max_samples: int = 32) -> List[Dict[str, Any]]: + """ + Convert PriorZero batch to LLM training samples. + + Args: + priorzero_batch: Tuple of (raw_obs_list, history_obs_list, llm_prior_per_tok_list, target_value, pred_value, cot_prefix_list, llm_action_list + CoT prefix list is added for CoT reuse optimization. + + Returns: + Tuple of (input_ids, attention_mask, action_mask, advantages, rollout_logprob) + """ + # Support both 7-element (legacy) and 8-element (with action_list) batch formats + if len(priorzero_batch) == 8: + raw_obs_list, history_obs_list, llm_prior_per_tok_list, target_value, pred_value, cot_prefix_list, llm_action_list, _action_list = priorzero_batch + else: + raw_obs_list, history_obs_list, llm_prior_per_tok_list, target_value, pred_value, cot_prefix_list, llm_action_list = priorzero_batch + + assert len(raw_obs_list) == len(history_obs_list) == len(llm_prior_per_tok_list) == len(target_value) == len(pred_value) == len(cot_prefix_list) == len(llm_action_list), \ + f"Batch size mismatch: raw_obs={len(raw_obs_list)}, history_obs={len(history_obs_list)}, llm_prior_per_tok={len(llm_prior_per_tok_list)}, \ + target_value={len(target_value)}, pred_value={len(pred_value)}, cot_prefix={len(cot_prefix_list)}, llm_action={len(llm_action_list)}" + + # Build samples with CoT prefixes + samples = self.build_llm_samples( + raw_obs_list, history_obs_list, llm_prior_per_tok_list, pred_value, target_value, cot_prefix_list, llm_action_list + ) + random.Random(0).shuffle(samples) + + + def _select_samples_with_unique_priority(sample_list, keep_n): + """优先取去重后的样本;如果去重后不够,则按原始顺序补齐。""" + if len(sample_list) < keep_n: + return None + unique_samples = unique_dicts_hash(sample_list) + if len(unique_samples) >= keep_n: + return unique_samples[:keep_n] + remain = keep_n - len(unique_samples) + selected = unique_samples + sample_list[:remain] + return selected[:keep_n] + + if ddp: + gathered_samples = [None for _ in range(self.world_size)] + dist.all_gather_object(gathered_samples, samples) + + global_samples = [] + for rank_samples in gathered_samples: + if rank_samples is not None: + global_samples.extend(rank_samples) + global_max_samples = self.world_size * max_samples + selected_global_samples = _select_samples_with_unique_priority(global_samples, global_max_samples) + + if selected_global_samples is None: + _log_train.warning( + f"Insufficient global samples: total={len(global_samples)} < required={global_max_samples}" + ) + return False, [global_samples] + + start = self.rank * max_samples + end = (self.rank + 1) * max_samples + real_samples = selected_global_samples[start:end] + _log_train.debug( + f"[Rank {self.rank}] local={len(samples)}, global={len(selected_global_samples)}, slice={start}:{end}" + ) + else: + selected_samples = _select_samples_with_unique_priority(samples, max_samples) + if selected_samples is None: + return False, [samples] + + per_rank = len(selected_samples) // self.world_size + start = self.rank * per_rank + end = (self.rank + 1) * per_rank if self.rank != self.world_size - 1 else len(selected_samples) + _log_train.debug(f"[Rank {self.rank}] samples slice={start}:{end}, total={len(selected_samples)}") + real_samples = selected_samples[start:end] + + if self.use_cot: + targets_only = [s["prefix_cot"] + " " + s["target"] + self.tokenizer.eos_token for s in real_samples] + if self.args.reward_func.format_reward: + fmt_rewards = torch.tensor([_format_reward(t) for t in targets_only]) + else: + fmt_rewards = None + else: + targets_only = ["Action: " +s["target"] + self.tokenizer.eos_token for s in real_samples] + fmt_rewards = None + + full_ids_list = [s['full_ids'] for s in real_samples] + tgt_ids_list = [s['label_ids'] for s in real_samples] + # Consistency check: decoded label_ids should match the expected target text. + # Convert hard assert to warning + filter to avoid crashing on tokenizer round-trip edge cases. + decoded_labels = self.tokenizer.batch_decode(tgt_ids_list) + if decoded_labels != targets_only: + mismatch_indices = [ + i for i, (d, t) in enumerate(zip(decoded_labels, targets_only)) if d != t + ] + _log_train.warning( + f"[make_llm_train_samples] label_ids decode mismatch for {len(mismatch_indices)}/{len(targets_only)} samples. " + f"First mismatch idx={mismatch_indices[0] if mismatch_indices else '?'}: " + f"decoded={decoded_labels[mismatch_indices[0]]!r:.120} vs expected={targets_only[mismatch_indices[0]]!r:.120}" + if mismatch_indices else "" + ) + # Filter out mismatched samples to avoid training on corrupted data + keep_mask = [i for i in range(len(targets_only)) if i not in set(mismatch_indices)] + if len(keep_mask) == 0: + _log_train.warning("[make_llm_train_samples] All samples mismatched, skipping batch") + return False, [real_samples] + real_samples = [real_samples[i] for i in keep_mask] + targets_only = [targets_only[i] for i in keep_mask] + full_ids_list = [full_ids_list[i] for i in keep_mask] + tgt_ids_list = [tgt_ids_list[i] for i in keep_mask] + if fmt_rewards is not None: + fmt_rewards = fmt_rewards[keep_mask] + inputs = self.tokenizer.pad({"input_ids": full_ids_list}, padding=True, return_tensors="pt") + labels = torch.full_like(inputs.input_ids, -100) + for i, tgt_ids in enumerate(tgt_ids_list): + tgt_len = len(tgt_ids) + labels[i, -tgt_len:] = inputs.input_ids[i, -tgt_len:] + action_mask_full = (labels != -100).long() + max_tgt_len = max(len(t) for t in tgt_ids_list) + action_mask = action_mask_full[:, -max_tgt_len:] + log_status_tmp = {} + log_status = [] + + if fmt_rewards is not None: + fmt_weight = self.args.reward_func.format_param.format_weight + assert 0.0 <= fmt_weight < 1.0, f"format_weight should be in [0, 1), but got {fmt_weight}" + log_status_tmp['fmt_rewards'] = fmt_rewards.tolist() + + # t 时刻的 target_value = td_step 步真实 r 的折扣和 + boostrap( t + td_step) 的 v + target_value = torch.tensor([s["target_value"] for s in real_samples], dtype=torch.float32) + # t 时刻的 pred_value = boostrap( t ) 的 v + pred_value = torch.tensor([s["pred_value"] for s in real_samples], dtype=torch.float32) + advantage = target_value - pred_value + + if self.args.advantage_type == "advantage": + advantage = advantage + log_status_tmp["value_advantage"] = advantage.tolist() + if fmt_rewards is not None: + advantage = (1 - fmt_weight) * advantage + fmt_weight * fmt_rewards + log_status_tmp["final_advantage"] = advantage.tolist() + + + elif self.args.advantage_type == "advantage_batch_norm": + # Legacy implementation: batch normalization (not recommended) + advantage = (advantage - advantage.mean()) / (advantage.std() + 1e-8) + log_status_tmp["value_advantage"] = advantage.tolist() + + if fmt_rewards is not None: + advantage = (1 - fmt_weight) * advantage + fmt_weight * fmt_rewards + log_status_tmp["final_advantage"] = advantage.tolist() + + elif self.args.advantage_type == "advantage_global_batch_norm": + # self.global_batch_advantages + self.global_batch_advantages += advantage.tolist() + advantage = (advantage - np.mean(self.global_batch_advantages)) / (np.std(self.global_batch_advantages) + 1e-8) + log_status_tmp["value_advantage"] = advantage.tolist() + + if fmt_rewards is not None: + advantage = (1 - fmt_weight) * advantage + fmt_weight * fmt_rewards + log_status_tmp["final_advantage"] = advantage.tolist() + elif self.args.advantage_type == "advantage_running_norm": + if self.value_normalizer is not None: + raw_mean = advantage.mean().item() + raw_std = advantage.std().item() + raw_min = advantage.min().item() + raw_max = advantage.max().item() + batch_size = advantage.numel() + + advantage, norm_stats = self.value_normalizer.normalize( + advantage, + clip_values=True, + return_stats=True + ) + + norm_min = advantage.min().item() + norm_max = advantage.max().item() + norm_mean = advantage.mean().item() + norm_std = advantage.std().item() + + if self.rank == 0 and self.value_normalizer.update_count % 10 == 0: + _log_train.debug( + f"[Value Norm] step={self.value_normalizer.update_count} | " + f"running: mean={norm_stats['running_mean']:.3f}, std={norm_stats['running_std']:.3f} | " + f"norm: min={norm_min:.3f}, max={norm_max:.3f}" + ) + else: + batch_mean = advantage.mean().item() + batch_std = advantage.std().item() + batch_min = advantage.min().item() + batch_max = advantage.max().item() + batch_size = advantage.numel() + + if self.value_count == 0: + self.value_running_mean = batch_mean + self.value_running_std = max(batch_std, 1e-8) # Avoid zero std + else: + self.value_running_mean = ( + self.running_momentum * self.value_running_mean + + (1 - self.running_momentum) * batch_mean + ) + self.value_running_std = ( + self.running_momentum * self.value_running_std + + (1 - self.running_momentum) * max(batch_std, 1e-8) + ) + + self.value_count += 1 + advantage = (advantage - self.value_running_mean) / (self.value_running_std + 1e-8) + + norm_min = advantage.min().item() + norm_max = advantage.max().item() + norm_mean = advantage.mean().item() + norm_std = advantage.std().item() + + if self.rank == 0 and self.value_count % 10 == 0: + _log_train.debug( + f"[Adv Norm] step={self.value_count} | " + f"running: mean={self.value_running_mean:.3f}, std={self.value_running_std:.3f} | " + f"norm: min={norm_min:.3f}, max={norm_max:.3f}" + ) + + log_status_tmp["value_advantage"] = advantage.tolist() + if fmt_rewards is not None: + advantage = (1 - fmt_weight) * advantage + fmt_weight * fmt_rewards + log_status_tmp["final_advantage"] = advantage.tolist() + else: + raise ValueError(f"Unknown advantage_type: {self.args.advantage_type}") + + log_status = [ + {k: log_status_tmp[k][i] for k in log_status_tmp.keys()} for i in range(len(log_status_tmp['value_advantage'])) + ] + + for i, s in enumerate(real_samples): + if len(s['rollout_logprob']) != len(s['label_ids']): + raise ValueError( + f"Length mismatch at sample {i}: " + f"len(rollout_logprob)={len(s['rollout_logprob'])}, " + f"len(label_ids)={len(s['label_ids'])}, " + f"target={repr(s['target'])}" + ) + old_seq_max_len = max([len(s['rollout_logprob']) for s in real_samples]) + rollout_logprob = torch.zeros(len(real_samples), old_seq_max_len, dtype=torch.float32) + for idx in range(len(real_samples)): + logprob_token_list = real_samples[idx]['rollout_logprob'] + rollout_logprob[idx, -len(logprob_token_list):] = torch.tensor(logprob_token_list, dtype=torch.float32) + + return True, (inputs.input_ids, inputs.attention_mask, action_mask, advantage, rollout_logprob, log_status) + + def _ensure_tp_pg(self) -> None: + """Lazy-init the vLLM TP subgroup PG (size = vllm_tensor_parallel_size). + + With `distributed_executor_backend='external_launcher'` and TP>1, vLLM partitions the + torch.dist world into TP subgroups. Each rank participates in exactly one subgroup of + consecutive ranks `[g_start, g_start + tp_size)`. We need that subgroup as a PG to + all_gather prompts so each TP partner submits the same input to `llm.generate()`. + + Note: `dist.new_group` is collective — every rank must call it for every group. + """ + if getattr(self, '_tp_pg', None) is not None: + return + tp_size = int(getattr(self.args, 'vllm_tensor_parallel_size', 1)) + rank = dist.get_rank() + for g_start in range(0, self.world_size, tp_size): + ranks = list(range(g_start, g_start + tp_size)) + pg = dist.new_group(ranks=ranks) + if rank in ranks: + self._tp_pg = pg + self._tp_group_start = g_start + + def _sync_prompts_for_tp(self, token_ids_list: List[List[int]]) -> Tuple[List[List[int]], slice]: + """Within a vLLM TP group (size > 1), all_gather the local prompt list and return + `(union, my_slice)`. Every rank in the TP group must submit the same `union` to + `vllm.generate` (so the V1 schedulers stay in lock-step on every rank), then slice + `outs[my_slice]` to recover its own outputs. + + For TP=1 / single-process / no DDP: returns `(list(token_ids_list), slice(0, n_real))`. + + Why count-only padding is insufficient: vLLM V1 runs an independent scheduler on each TP + rank from the same logical inputs, and any per-prompt length difference produces a + different chunked-prefill batch shape, which then hits a mismatch in the TP all_gather + of logits. Content must match, not just count. + """ + n_real = len(token_ids_list) + tp_size = int(getattr(self.args, 'vllm_tensor_parallel_size', 1)) + if not (dist.is_initialized() and self.world_size > 1 and tp_size > 1): + return list(token_ids_list), slice(0, n_real) + self._ensure_tp_pg() + rank = dist.get_rank() + local_idx = rank - self._tp_group_start + gathered: List[Optional[List[List[int]]]] = [None] * tp_size + dist.all_gather_object(gathered, list(token_ids_list), group=self._tp_pg) + offsets = [0] + for sub in gathered: + offsets.append(offsets[-1] + len(sub)) + union: List[List[int]] = [] + for sub in gathered: + union.extend(sub) + return union, slice(offsets[local_idx], offsets[local_idx + 1]) + + @torch.no_grad() + def drain_vllm_iter(self) -> None: + """Match the two `vllm.generate` calls that `get_llm_prior` would make in one outer eval + iter, but submit only the partners' prompts (via `_sync_prompts_for_tp([])`). Used by + the evaluator in DDP drain mode so TP partners can finish their real calls. No-op for + TP=1 / single-process. + """ + tp_size = int(getattr(self.args, 'vllm_tensor_parallel_size', 1)) + if not (dist.is_initialized() and self.world_size > 1 and tp_size > 1): + return + + cot_sampling_params = SamplingParams( + temperature=1.0, + top_p=1.0, + max_tokens=self.generate_max_len, + stop=["\n\n"], + include_stop_str_in_output=True, + logprobs=None, + prompt_logprobs=None, + ) + union, _ = self._sync_prompts_for_tp([]) + if union: + self.vllm_engine.add_requests(sampling_params=cot_sampling_params, prompt_token_ids=union) + self.vllm_engine.get_responses() + + score_sampling_params = SamplingParams( + temperature=self.temperature, + top_p=self.top_p, + max_tokens=1, + include_stop_str_in_output=True, + logprobs=None, + prompt_logprobs=1, + ) + union, _ = self._sync_prompts_for_tp([]) + if union: + self.vllm_engine.add_requests(sampling_params=score_sampling_params, prompt_token_ids=union) + self.vllm_engine.get_responses() + + @torch.no_grad() + def _build_cot_prefix_texts(self, all_user_prompts: List[str]) -> List[str]: + """ + 生成CoT推理前缀。 + 优化: 使用较短的max_tokens(128)和stop条件以减少不必要的生成。 + 从最后一次出现的 "Action:" 截断出 prefix(包含 Action: 和其后的空格位置)。 + 返回 prefix_cot_list,与 all_user_prompts 等长。 + """ + cot_sampling_params = SamplingParams( + temperature=1.0, + top_p=1.0, + max_tokens=self.generate_max_len, + stop=["\n\n"], + # stop=["Action:", "\n\n"] + include_stop_str_in_output=True, + logprobs=None, + prompt_logprobs=None, + ) + + all_context_texts = [self.build_chat_context(p) for p in all_user_prompts] + context_token_ids = self.tokenizer( + all_context_texts, + add_special_tokens=False, + max_length=self.prompt_max_len, + padding=False, + truncation=True, + )["input_ids"] + + context_token_ids_union, my_slice_cot = self._sync_prompts_for_tp(context_token_ids) + + self.vllm_engine.add_requests(sampling_params=cot_sampling_params, prompt_token_ids=context_token_ids_union) + cot_outputs = self.vllm_engine.get_responses() + cot_outputs = cot_outputs[my_slice_cot] + + prefix_cot_list, full_output = [], [] + reasoning_pattern = re.compile(r"Reasoning\s*:", re.IGNORECASE) + action_pattern = re.compile(r"Action\s*:", re.IGNORECASE) + + for output in cot_outputs: + gen_text = output.outputs[0].text + full_output.append(gen_text) + # TODO 这里是否要清洗数据?清洗过后,计算prior先验的时候比较正常,但是format_reward几乎没用 + # if not reasoning_pattern.search(gen_text): + # prefix_cot_list.append("Action:") + # continue + action_match = action_pattern.search(gen_text) + if action_match: + end_index = action_match.end() + prefix_piece = gen_text[:end_index].strip() + else: + # prefix_piece = gen_text.strip() + prefix_piece = gen_text.strip() + "\nAction:" + + prefix_cot_list.append(prefix_piece) + + return prefix_cot_list, full_output + + @torch.no_grad() + def get_llm_prior( + self, + states: List[str], + valid_actions_list: List[List[str]], + histories: Optional[List[List[Tuple[str, str, float]]]] = None, + return_cot: bool = False, # CoT reuse optimization: return CoT prefixes + ) -> List[Any]: + """ + Get LLM prior scores for actions. + + Args: + states: List of current state observations + valid_actions_list: List of valid actions for each state + histories: List of history observations + return_cot: If True, return CoT prefixes for reuse (optimization) + + Returns: + If return_cot=False: (llm_prior_per_seq, llm_prior_per_tok) + If return_cot=True: (llm_prior_per_seq, llm_prior_per_tok, prefix_cots, full_cot_outputs) + """ + prompt_list = [] + assert len(states) == len(histories) == len(valid_actions_list) + for state, history, valid_actions in zip(states, histories, valid_actions_list): + prompt = self.get_user_prompt(current_obs=state, history=history, valid_actions=valid_actions) + prompt_list.append(prompt) + + if self.use_cot: + prefix_cots, full_output = self._build_cot_prefix_texts(prompt_list) + else: + prefix_cots = [None] * len(prompt_list) + full_output = None + + all_prompts = [] + all_labels = [] + all_prefix_cots = [] + all_env_indices = [] + + for env_idx, (prompt, actions, prefix) in enumerate(zip(prompt_list, valid_actions_list, prefix_cots)): + actions2 = actions if "go" in actions else (actions + ["go"]) # 确保环境使用的动作都在valid actions里有对应的logprob + for action in actions2: + all_prompts.append(prompt) + all_labels.append(action) + all_prefix_cots.append(prefix) + all_env_indices.append(env_idx) + assert len(all_prompts) == len(all_labels) == len(all_prefix_cots) == len(all_env_indices) + + scores, rollout_action_logprob, full_ids, label_ids = self._score_labels_with_prompt_logprobs(all_prompts, all_labels, all_prefix_cots) + assert len(all_prompts) == len(scores) == len(rollout_action_logprob) == len(full_ids) == len(label_ids) + + llm_prior_per_seq, llm_prior_per_tok = [],[], + cur_env_idx = 0 + seq_dict = {} + tok_dict = {'rollout_action_logprob': {}, 'full_ids': {}, 'label_ids': {}} + + for idx, (env_idx, prompt, label, prefix_cot) in enumerate(zip(all_env_indices, all_prompts, all_labels, all_prefix_cots)): + if env_idx != cur_env_idx: + llm_prior_per_seq.append(seq_dict) + llm_prior_per_tok.append(tok_dict) + seq_dict = {} + tok_dict = {'rollout_action_logprob': {}, 'full_ids': {}, 'label_ids': {}} + cur_env_idx = env_idx + + seq_dict[label] = scores[idx] + tok_dict['rollout_action_logprob'][label] = rollout_action_logprob[idx] + tok_dict['full_ids'][label] = full_ids[idx] + tok_dict['label_ids'][label] = label_ids[idx] + tok_dict['prompt'] = prompt + tok_dict['prefix_cot'] = prefix_cot + tok_dict['current_obs'] = states[env_idx] + tok_dict['history'] = histories[env_idx] + + if len(seq_dict) > 0: + llm_prior_per_seq.append(seq_dict) + llm_prior_per_tok.append(tok_dict) + + # Drain mode (empty inputs from caller): prompt_list / llm_prior_per_seq are empty, + # so skip the per-call episode log to avoid IndexError on prompt_list[0]. + if len(prompt_list) > 0 and len(llm_prior_per_seq) > 0: + self.episode_output.append({ + "Instruction": prompt_list[0], + "Response": full_output[0] if full_output else "(no CoT)", + "llm_prior_per_seq": llm_prior_per_seq[0] + }) + # CoT reuse optimization: return CoT prefixes if requested + if return_cot: + return llm_prior_per_seq, llm_prior_per_tok, prefix_cots, full_output + else: + return llm_prior_per_seq, llm_prior_per_tok + + @torch.no_grad() + def _score_labels_with_prompt_logprobs(self, all_prompts: List[str], all_labels: List[str], all_prefix_cots: List[str]) -> List[float]: + assert len(all_prompts) == len(all_labels) == len(all_prefix_cots) + sampling_params = SamplingParams( + temperature=self.temperature, + top_p=self.top_p, + max_tokens=1, + include_stop_str_in_output=True, + logprobs=None, + prompt_logprobs=1, + ) + + all_context_texts = [self.build_chat_context(p) for p in all_prompts] + context_ids = self.tokenizer(all_context_texts, add_special_tokens=False, max_length=self.prompt_max_len - self.generate_max_len - 20, padding=False, truncation=True)["input_ids"] + + if self.use_cot: + label_texts = [pc + " " + l + self.tokenizer.eos_token for pc, l in zip(all_prefix_cots, all_labels)] + label_texts_no_cots = [" " + l + self.tokenizer.eos_token for l in all_labels] + else: + label_texts = ["Action: " + l + self.tokenizer.eos_token for l in all_labels] + label_texts_no_cots = label_texts + + label_ids = self.tokenizer(label_texts, add_special_tokens=False, padding=False, truncation=False)["input_ids"] + label_ids_no_cots = self.tokenizer(label_texts_no_cots, add_special_tokens=False, padding=False, truncation=False)["input_ids"] + + for idx, (l_ids, l_ids_not_cot) in enumerate(zip(label_ids, label_ids_no_cots)): + len_not_cot = len(l_ids_not_cot) + if l_ids[-len_not_cot:] != l_ids_not_cot: + raise ValueError(f"Label IDs mismatch: with CoT {l_ids[-len_not_cot:]}, without CoT {l_ids_not_cot}, label_text: {label_texts[idx]}") + + full_ids = [c + l for c, l in zip(context_ids, label_ids)] + p_lens = [len(x) for x in context_ids] + l_lens = [len(x) for x in label_ids] + l_no_cots_lens = [len(x) for x in label_ids_no_cots] + + full_ids_union, my_slice_score = self._sync_prompts_for_tp(full_ids) + + self.vllm_engine.add_requests(sampling_params=sampling_params, prompt_token_ids=full_ids_union) + outs = self.vllm_engine.get_responses() + outs = outs[my_slice_score] + + scores = [] + rollout_action_logprob = [] + nan_found = False + for i, (out, ids, p_len, l_len, l_no_cots_len) in enumerate(zip(outs, full_ids, p_lens, l_lens, l_no_cots_lens)): + prompt_logprobs = getattr(out, "prompt_logprobs", None) + token_lps = [] + + for j in range(1, len(ids)): + tok_id = ids[j] + lp_dict = prompt_logprobs[j] + + assert tok_id in lp_dict + token_lps.append(lp_dict[tok_id].logprob) + + if not token_lps: + scores.append(float("-inf")) + rollout_action_logprob.append([]) + else: + assert l_no_cots_len <= l_len + if self.llm_prior_with_cot: + target_lps = token_lps[-l_len:] + else: + target_lps = token_lps[-l_no_cots_len:] + denom = len(target_lps) + + score = sum(target_lps) if self.reduction == "sum" else sum(target_lps) / denom + scores.append(score) + + if (not nan_found) and math.isnan(score): + vllm_returned_nan = any(math.isnan(x) for x in target_lps) + token_level_debug = [] + for t_id, t_lp in zip(ids[1:], token_lps): + token_level_debug.append(f"TokenID: {t_id} -> LogProb: {t_lp} {'(NaN HERE!)' if math.isnan(t_lp) else ''}") + + nan_found = True + nan_debug_dump = ( + f"\n{'='*20} [NaN DEBUG REPORT] {'='*20}\n" + f"Sample Index (i): {i}\n" + f"Reason: {'vLLM returned NaN logprob' if vllm_returned_nan else 'Math error during sum/div'}\n\n" + f"--- Text Info ---\n" + f"Prompt: ...{repr(all_prompts[i])}\n" + f"Label Action: {repr(all_labels[i])}\n" + f"Prefix CoT: {repr(all_prefix_cots[i])}\n\n" + f"--- Numerical Info (Copy this to reproduce) ---\n" + f"Full Input Token IDs (full_ids[{i}]): {ids}\n" + f"Context Length (p_len): {p_len}\n" + f"Label Length (l_len): {l_len}\n" + f"Target Length (l_no_cots_len): {l_no_cots_len}\n\n" + f"--- Critical Calculation Data ---\n" + f"Head 10 Token IDs: {ids[1:11]}\n" + f"LogProbs List: {token_lps[:10]}\n" + f"Detailed Mapping:\n" + "\n".join(token_level_debug[:10]) + "\n\n" + + f"Tail Token IDs: {ids[-l_len - 10: -l_len]}\n" + f"LogProbs List: {token_lps[-l_len - 10: -l_len]}\n" + f"Detailed Mapping:\n" + "\n".join(token_level_debug[-l_len - 10: -l_len]) + "\n\n" + + f"Target Token IDs: {ids[-l_no_cots_len:]}\n" + f"LogProbs List: {target_lps}\n" + f"Detailed Mapping:\n" + "\n".join(token_level_debug[-l_no_cots_len:]) + "\n" + f"{'='*60}\n" + ) + rollout_action_logprob.append(token_lps[-l_len:]) + + if self.rank == 0: + if nan_found: + self._logger.info(nan_debug_dump) + + return scores, rollout_action_logprob, full_ids, label_ids + + @torch.no_grad() + def get_llm_output_log(self, wm_train_iter: int = 0, llm_train_iter: int = 0): + if self.rank != 0: + return + + self._logger.info( + f"\n{'='*80}\n" + f"[LLM Output Log] WM Iter: {wm_train_iter} | LLM Iter: {llm_train_iter}\n" + f"{'='*80}" + ) + + for i, tmp_dict in enumerate(self.episode_output[:15]): + instruction = tmp_dict["Instruction"] + response = tmp_dict["Response"] + llm_prior = tmp_dict["llm_prior_per_seq"] + + self._logger.info( + f"\n{'-'*80}\n" + f"[Step {i}]\n" + f"{'-'*80}\n" + f"Instruction:\n{instruction}\n\n" + f"Response:\n{response}\n\n" + f"Action Probabilities:" + ) + + action_probs = {a: math.exp(float(lp)) for a, lp in llm_prior.items() if lp is not None and math.isfinite(float(lp))} + all_prob = sum(action_probs.values()) + + for action, prob in sorted(action_probs.items(), key=lambda x: x[1], reverse=True): + self._logger.info(f" {action:30s} | unnorm={prob:.6f} | norm={(prob / all_prob):.6f}") + self._logger.info(f" {'':30s} | unnorm={1-all_prob:.6f}") + self.episode_output = [] + + + def clear_statis(self): + if self.value_normalizer is not None: + self.value_normalizer.clear() + self.global_batch_advantages.clear() + \ No newline at end of file diff --git a/zoo/jericho/priorzero/src/priorzero_entry_sync.py b/zoo/jericho/priorzero/src/priorzero_entry_sync.py new file mode 100644 index 000000000..08a45c335 --- /dev/null +++ b/zoo/jericho/priorzero/src/priorzero_entry_sync.py @@ -0,0 +1,376 @@ +import sys +import os +from pathlib import Path + +import asyncio +import os +import sys +from functools import partial +from pathlib import Path +from typing import Tuple, Optional + +import torch +import torch.distributed as dist +import wandb + +from ding.config import compile_config, save_config +from ding.envs import create_env_manager, get_vec_env_setting +from ding.policy import create_policy +from ding.utils import set_pkg_seed, get_rank, get_world_size +from ding.worker import create_buffer, BaseLearner +from tensorboardX import SummaryWriter +from loguru import logger +import deepspeed + +from priorzero_config import ( + get_priorzero_config, + get_priorzero_debug_config, + get_available_models, +) +from priorzero_collector import PriorZeroCollector +from priorzero_evaluator import PriorZeroEvaluator +from priorzero_policy import * +from lzero.mcts.buffer.game_buffer_priorzero import PriorZeroGameBufferOptimized +from utils import dump_dataclass_cfg_py + +from lzero.entry.utils import calculate_update_per_collect + +def prepare_unizero(rank, cfg, create_cfg, llm_cfg, seed): + cfg = compile_config(cfg, seed=seed, auto=True, create_cfg=create_cfg) + env_fn, collector_env_cfg, evaluator_env_cfg = get_vec_env_setting(cfg.env) + collector_env = create_env_manager(cfg.env.manager, [partial(env_fn, cfg=c) for c in collector_env_cfg]) + evaluator_env = create_env_manager(cfg.env.manager, [partial(env_fn, cfg=c) for c in evaluator_env_cfg]) + + collector_env.seed(seed) + evaluator_env.seed(seed, dynamic_seed=False) + + policy = create_policy( cfg.policy, enable_field=['learn', 'collect', 'eval'], exp_name=cfg.exp_name, llm_cfg=llm_cfg) + logger.info(f"[Rank {rank}] Policy created") + + if cfg.policy.model_path is not None: + logging.info(f"Loading pretrained model from {cfg.policy.model_path}...") + policy.learn_mode.load_state_dict(torch.load(cfg.policy.model_path, map_location=cfg.policy.device)) + + os.makedirs(f'./{cfg.exp_name}/log/', exist_ok=True) + tb_logger = SummaryWriter(os.path.join(f'./{cfg.exp_name}/log/', 'serial')) if get_rank() == 0 else None + logger.info(f"[Rank {rank}] TensorBoard logger: ./{cfg.exp_name}/log/") + + learner = BaseLearner( + cfg.policy.learn.learner, + policy.learn_mode, + tb_logger, + exp_name=cfg.exp_name + ) + logger.info(f"[Rank {rank}] BaseLearner created") + + + replay_buffer = PriorZeroGameBufferOptimized(cfg.policy) + logger.info(f"[Rank {rank}] PriorZero replay buffer created (with game_segments support)") + + # Create collector + collector = PriorZeroCollector( + env=collector_env, + policy=policy.collect_mode, + llm_config=llm_cfg, + tb_logger=tb_logger, + exp_name=cfg.exp_name, + policy_config=cfg.policy, + ) + logger.info(f"[Rank {rank}] Collector created") + + # Create evaluator + evaluator = PriorZeroEvaluator( + n_evaluator_episode=cfg.env.n_evaluator_episode, + stop_value=cfg.env.stop_value, + env=evaluator_env, + policy=policy.eval_mode, + tb_logger=tb_logger, + exp_name=cfg.exp_name, + policy_config=cfg.policy, + llm_config=llm_cfg, + ) + logger.info(f"[Rank {rank}] Evaluator created") + learner.call_hook('before_run') + + return cfg, replay_buffer, tb_logger, policy, collector, evaluator, learner + +def bcast_obj(world_size, obj, rank, src=0): + if world_size <= 1: + return obj + lst = [obj] if rank == src else [None] + dist.broadcast_object_list(lst, src=src) + return lst[0] + +def train_priorzero( + cfg: dict, + create_cfg: dict, + llm_cfg, + seed: int = 0, + max_train_iter: int = int(1e6), + max_env_step: Optional[int] = int(1e10), + enable_profile: bool = False +): + rank = int(os.environ.get("RANK", "0")) + print(f"rank={rank}") + if rank == 0: + cfg, replay_buffer, tb_logger, policy, collector, evaluator, learner = prepare_unizero( + rank=rank, + cfg=cfg, + create_cfg=create_cfg, + llm_cfg=llm_cfg, + seed=seed) + batch_size = cfg.policy.batch_size + logger.info(f"[Rank {rank}] World Model components initialized") + dump_dataclass_cfg_py(llm_cfg, path=f"{cfg.exp_name}/llm_cfg.py") + llm_cfg.save_path = f'./{cfg.exp_name}/llm_ckpt/' + + from utils import Profiler + prof = Profiler(log_interval=10, stats_file=f'./{cfg.exp_name}/log/profiler.txt', enable_profile=enable_profile) + + from strategy.deepspeed import get_strategy, torch_dist_barrier_and_cuda_sync + strategy = get_strategy(llm_cfg) + strategy.print(llm_cfg) + + strategy.setup_distributed() # torchrun 下:绑定 local_rank + init_distributed + world_size = getattr(strategy, "world_size", 1) + + logger.info(f"[Rank {rank}] Initializing LLM Actor...") + set_pkg_seed(seed + rank, use_cuda=True) + + from models.actor import PolicyModel, ReferenceModel + if llm_cfg.rft_kl_coef > 0: + ref_model = ReferenceModel( + strategy=strategy, + pretrain=llm_cfg.model_name_or_path + ) + else: + ref_model = None + + from vllm_utils.vllm_engine import create_vllm_engine + vllm_engine = create_vllm_engine( + tensor_parallel_size=llm_cfg.vllm_tensor_parallel_size, + pretrain=llm_cfg.model_name_or_path, + enable_prefix_caching=llm_cfg.enable_prefix_caching, + max_model_len=llm_cfg.prompt_max_len + llm_cfg.generate_max_len, + gpu_memory_utilization=llm_cfg.gpu_memory_utilization, + vllm_enable_sleep=llm_cfg.vllm_enable_sleep, + ) + + print(f'[Rank {rank}] Vllm engine successfully created!') + + from priorzero_datafactory import DataProcessor + data_processor = DataProcessor(rank=rank, + world_size=world_size, + vllm_engine=vllm_engine, + strategy=strategy, + model_path=llm_cfg.model_name_or_path, + exp_name=cfg.exp_name if rank == 0 else None, + ) + if rank == 0: + collector.data_processor = data_processor + collector.prof = prof + evaluator.data_processor = data_processor + + policy_model = PolicyModel( + strategy=strategy, + pretrain=llm_cfg.model_name_or_path, + vllm_engine=vllm_engine, + max_steps=llm_cfg.max_steps + ) + from priorzero_trainer import PriorZeroLLMTrainer + trainer = PriorZeroLLMTrainer( + cfg=llm_cfg, + pretrain=llm_cfg.model_name_or_path, + strategy= strategy, + vllm_engine = vllm_engine, + policy_model=policy_model, + reference_model=ref_model, + exp_name=cfg.exp_name if rank == 0 else None, + tb_logger=tb_logger if rank == 0 else None, + llm_save_freq=llm_cfg.llm_save_freq + ) + + torch_dist_barrier_and_cuda_sync() + + train_schedule = llm_cfg.train_schedule + train_alternate = train_schedule["alternate"] + current_phase = None + llm_collect_mode = None + if train_alternate: + current_phase = train_schedule["start_phase"] + last_wm_train_iter = 0 + last_llm_train_iter = 0 + llm_collect_mode = train_schedule["llm_collect_mode"] + + while True: + cmd = "noop" + priorzero_batch = None + if rank == 0: + if learner.train_iter != 0 and evaluator.should_eval(wm_train_iter=learner.train_iter, llm_train_iter=policy_model.train_iter, phase=current_phase): + logger.info(f"\n[Rank {rank}: Iter {learner.train_iter}] Evaluating...") + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.wake_up() + evaluator.eval(wm_train_iter=learner.train_iter, llm_train_iter=policy_model.train_iter, phase=current_phase) + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.sleep() + + if cmd != "stop": + if not train_alternate or (train_alternate and current_phase == "wm") or (train_alternate and current_phase == "llm" and llm_collect_mode != "no_collect"): + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.wake_up() + + new_data = collector.collect(train_iter=learner.train_iter, policy_kwargs={'temperature': 0.25, 'epsilon': 0.0}, phase=current_phase) + data_processor.get_llm_output_log(wm_train_iter=learner.train_iter, llm_train_iter=policy_model.train_iter) + + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.sleep() + + update_per_collect = calculate_update_per_collect(cfg, new_data, world_size=1) + + replay_buffer.push_game_segments(new_data) + replay_buffer.remove_oldest_data_to_fit() + + num_of_transitions = replay_buffer.get_num_of_transitions() + new_num_of_transitions = replay_buffer.get_num_of_transitions() - replay_buffer.last_pos_in_transition + logger.info(f"[Rank {rank}] Data collected, num_of_transitions: {num_of_transitions} transitions\tnew_num_of_transitions: {new_num_of_transitions}") + + if not (num_of_transitions > batch_size): + logger.warning( + f' ⚠ Data in replay_buffer is not sufficient: ' + f'batch_size: {batch_size}, replay_buffer: {replay_buffer}. Continue to collect...' + ) + cmd = "noop" + cmd = bcast_obj(world_size, cmd, rank, src=0) + continue + + if llm_cfg.enable_world_model and (not train_alternate or (train_alternate and current_phase == "wm")): + logger.info(f"[Rank {rank}: World Model] [Iter {learner.train_iter}] Training for {update_per_collect} updates......") + for i in range(update_per_collect): + with prof.block("train_world_model", rank=0): + train_data = replay_buffer.sample(batch_size, policy) + train_data.append(learner.train_iter) + + log_vars = learner.train(train_data, collector.envstep) + if cfg.policy.use_priority: + replay_buffer.update_priority(train_data, log_vars[0]['value_priority_orig']) + policy.recompute_pos_emb_diff_and_clear_cache() + if llm_cfg.enable_rft and train_alternate and learner.train_iter - last_wm_train_iter >= train_schedule["wm_update_iters"]: + current_phase = "llm" + last_wm_train_iter = learner.train_iter + if llm_collect_mode != "no_collect": + replay_buffer.mark_latest_transitions_consumed() + continue + + if llm_cfg.enable_rft and (not train_alternate or (train_alternate and current_phase == "llm")): + print(f"[Rank 0] world_model: train_iter ={learner.train_iter} \t replay_buffer.fetch_latest_batch begin \t") + if llm_collect_mode != "no_collect": + priorzero_batch = replay_buffer.fetch_latest_batch(batch_size=-1, policy=policy, select_last=True) + else: + priorzero_batch = replay_buffer.fetch_latest_batch(batch_size=128, policy=policy, select_last=False) + # 清理 policy的cahce,防止OOM + torch.cuda.empty_cache() + print(f"[Rank 0] fetch_latest_batch returned: type={type(priorzero_batch)}, len={len(priorzero_batch)}") + cmd = "llm" + + if collector.envstep >= max_env_step or learner.train_iter >= max_train_iter: + cmd = "stop" + + cmd = bcast_obj(world_size, cmd, rank, src=0) + if cmd == "stop": + break + elif cmd == "llm": + with prof.block("train_llm", rank=rank): + logger.info(f"[Rank {rank}] Waiting for broadcast of train_samples from Rank 0...") + priorzero_batch = bcast_obj(world_size, priorzero_batch, rank, src=0) + logger.info(f"[Rank {rank}] Received broadcast. train_samples count: {len(priorzero_batch[0]) if priorzero_batch and len(priorzero_batch) > 0 else 'UNKNOWN'}. Starting LLM training...") + + llm_need_sample_cnt = llm_cfg.train_batch_size * llm_cfg.max_rollout_staleness // 1 + flag, train_samples = data_processor.make_llm_train_samples(priorzero_batch, max_samples=llm_need_sample_cnt) + if not flag: # 检查样本是否有效 + logger.warning(f"[Rank {rank}] No valid LLM training samples were created. Skipping this LLM training phase.") + continue + + trainer.train_batch(train_samples, collect_env_steps=collector.envstep) + if llm_collect_mode != "no_collect": + replay_buffer.mark_latest_transitions_consumed() + torch_dist_barrier_and_cuda_sync() + + if llm_cfg.enable_world_model and train_alternate and trainer.global_step - last_llm_train_iter >= train_schedule["llm_update_iters"]: + current_phase = "wm" + last_llm_train_iter = trainer.global_step + data_processor.clear_statis() + + +def main(): + """ + Main entry point with argument parsing. + """ + import argparse + + parser = argparse.ArgumentParser( + description='PriorZero Training with Auto Model Configuration', + formatter_class=argparse.RawDescriptionHelpFormatter, + epilog=""" +Examples: + # Use default model (qwen2.5-1.5b) + torchrun --nproc_per_node 2 priorzero_entry_sync.py + + # Use specific model + torchrun --nproc_per_node 2 priorzero_entry_sync.py --model qwen2.5-0.5b + torchrun --nproc_per_node 2 priorzero_entry_sync.py --model qwen2.5-7b + + # List all available models + python priorzero_entry_sync.py --list-models + + # Different environment + torchrun --nproc_per_node 2 priorzero_entry_sync.py --env_id zork1.z5 --model qwen2.5-1.5b + """ + ) + parser.add_argument('--env_id', type=str, default='detective.z5', help='Jericho game ID') + parser.add_argument('--seed', type=int, default=0, help='Random seed') + parser.add_argument('--max_iter', type=int, default=int(1e6), help='Max training iterations') + parser.add_argument('--quick_test', action='store_true', default=False, help='Use quick test config') + # Model selection + parser.add_argument('--model', type=str, default="qwen2.5-3b", choices=get_available_models()) + parser.add_argument('--enable_profile', action='store_true', default=False) + parser.add_argument('--use_cot', action='store_true', default=False) + args = parser.parse_args() + + model_key = args.model if args.model else "qwen2.5-1.5b" + print(f"\n{'='*80}") + print(f"PriorZero Training Configuration") + print(f"{'='*80}") + print(f"Environment: {args.env_id}") + print(f"Model: {model_key}") + print(f"Seed: {args.seed}") + print(f"Quick Test: {args.quick_test}") + print(f"use cot: {args.use_cot}") + print(f"enable_profile: {args.enable_profile}") + print(f"{'='*80}\n") + + if args.quick_test: + logger.info("Using quick test configuration") + main_cfg, create_cfg, llm_cfg = get_priorzero_debug_config( + args.env_id, args.seed, use_cot=args.use_cot, + exp_name=f'data_priorzero/priorzero_debug_{args.env_id}', + model_key=model_key, + ) + else: + main_cfg, create_cfg, llm_cfg = get_priorzero_config( + args.env_id, args.seed, use_cot=args.use_cot, + model_key=model_key, + ) + + train_priorzero( + main_cfg, + create_cfg, + llm_cfg, + seed=args.seed, + max_train_iter=args.max_iter, + enable_profile=args.enable_profile, # 是否要对各个耗时部分进行 profile + ) + + +if __name__ == "__main__": + os.environ['TOKENIZERS_PARALLELISM'] = 'false' + main() diff --git a/zoo/jericho/priorzero/src/priorzero_entry_sync_ddp.py b/zoo/jericho/priorzero/src/priorzero_entry_sync_ddp.py new file mode 100644 index 000000000..76f3f7a04 --- /dev/null +++ b/zoo/jericho/priorzero/src/priorzero_entry_sync_ddp.py @@ -0,0 +1,379 @@ +import sys +import os +from pathlib import Path + +import asyncio +import os +import sys +from functools import partial +from pathlib import Path +from typing import Tuple, Optional + +import torch +import torch.distributed as dist +import wandb + +from ding.config import compile_config, save_config +from ding.envs import create_env_manager, get_vec_env_setting +from ding.policy import create_policy +from ding.utils import set_pkg_seed, get_rank, get_world_size +from ding.worker import create_buffer, BaseLearner +from tensorboardX import SummaryWriter +from loguru import logger +import deepspeed + +from priorzero_config import ( + get_priorzero_config, + get_priorzero_debug_config, + get_available_models, +) +from priorzero_collector import PriorZeroCollector +from priorzero_evaluator import PriorZeroEvaluator +from priorzero_policy import * +from lzero.mcts.buffer.game_buffer_priorzero import PriorZeroGameBufferOptimized +from utils import dump_dataclass_cfg_py + +from lzero.entry.utils import calculate_update_per_collect + +def prepare_unizero(rank, cfg, create_cfg, llm_cfg, seed): + cfg = compile_config(cfg, seed=seed, auto=True, create_cfg=create_cfg) + env_fn, collector_env_cfg, evaluator_env_cfg = get_vec_env_setting(cfg.env) + collector_env = create_env_manager(cfg.env.manager, [partial(env_fn, cfg=c) for c in collector_env_cfg]) + evaluator_env = create_env_manager(cfg.env.manager, [partial(env_fn, cfg=c) for c in evaluator_env_cfg]) + + collector_env.seed(seed) + evaluator_env.seed(seed, dynamic_seed=False) + + policy = create_policy( cfg.policy, enable_field=['learn', 'collect', 'eval'], exp_name=cfg.exp_name, llm_cfg=llm_cfg) + if cfg.policy.model_path is not None: + logging.info(f"[Rank {rank}] Loading pretrained model from {cfg.policy.model_path}...") + policy.learn_mode.load_state_dict(torch.load(cfg.policy.model_path, map_location=cfg.policy.device)) + logger.info(f"[Rank {rank}] Policy created") + + os.makedirs(f'./{cfg.exp_name}/log/', exist_ok=True) + tb_logger = SummaryWriter(os.path.join(f'./{cfg.exp_name}/log/', 'serial')) if get_rank() == 0 else None + logger.info(f"[Rank {rank}] TensorBoard logger: ./{cfg.exp_name}/log/") + + learner = BaseLearner( + cfg.policy.learn.learner, + policy.learn_mode, + tb_logger, + exp_name=cfg.exp_name + ) + logger.info(f"[Rank {rank}] BaseLearner created") + + + replay_buffer = PriorZeroGameBufferOptimized(cfg.policy) + logger.info(f"[Rank {rank}] PriorZero replay buffer created (with game_segments support)") + + # Create collector + collector = PriorZeroCollector( + env=collector_env, + policy=policy.collect_mode, + llm_config=llm_cfg, + tb_logger=tb_logger, + exp_name=cfg.exp_name, + policy_config=cfg.policy, + ) + logger.info(f"[Rank {rank}] Collector created") + + # Create evaluator + evaluator = PriorZeroEvaluator( + n_evaluator_episode=cfg.env.n_evaluator_episode, + stop_value=cfg.env.stop_value, + env=evaluator_env, + policy=policy.eval_mode, + tb_logger=tb_logger, + exp_name=cfg.exp_name, + policy_config=cfg.policy, + llm_config=llm_cfg, + ) + logger.info(f"[Rank {rank}] Evaluator created") + learner.call_hook('before_run') + + return cfg, replay_buffer, tb_logger, policy, collector, evaluator, learner + +def all_gather_cmd(world_size, obj) -> List: + if world_size <= 1: + return [obj] + lst = [None] * dist.get_world_size() + dist.all_gather_object(lst, obj) + return lst + +def train_priorzero( + cfg: dict, + create_cfg: dict, + llm_cfg, + seed: int = 0, + max_train_iter: int = int(1e6), + max_env_step: Optional[int] = int(1e10), + enable_profile: bool = False +): + rank = int(os.environ.get("RANK", "0")) + print(f"DEBUG: Is dist initialized at start? {dist.is_initialized()}") + if dist.is_initialized(): + print(f"DEBUG: Backend is {dist.get_backend()}") + from strategy.deepspeed import get_strategy, torch_dist_barrier_and_cuda_sync + strategy = get_strategy(llm_cfg) + strategy.print(llm_cfg) + + strategy.setup_distributed() # torchrun 下:绑定 local_rank + init_distributed + world_size = getattr(strategy, "world_size", 1) + + + cfg, replay_buffer, tb_logger, policy, collector, evaluator, learner = prepare_unizero( + rank=rank, + cfg=cfg, + create_cfg=create_cfg, + llm_cfg=llm_cfg, + seed=seed) + batch_size = cfg.policy.batch_size + logger.info(f"[Rank {rank}] World Model components initialized") + if rank == 0: + dump_dataclass_cfg_py(llm_cfg, path=f"{cfg.exp_name}/llm_cfg.py") + llm_cfg.save_path = f'./{cfg.exp_name}/llm_ckpt/' + + from utils import Profiler + prof = Profiler(log_interval=10, stats_file=f'./{cfg.exp_name}/log/profiler.txt', enable_profile=enable_profile) + + + logger.info(f"[Rank {rank}] Initializing LLM Actor...") + set_pkg_seed(seed + rank, use_cuda=True) + + from models.actor import PolicyModel, ReferenceModel + if llm_cfg.rft_kl_coef > 0: + ref_model = ReferenceModel( + strategy=strategy, + pretrain=llm_cfg.model_name_or_path + ) + else: + ref_model = None + + from vllm_utils.vllm_engine import create_vllm_engine + vllm_engine = create_vllm_engine( + tensor_parallel_size=llm_cfg.vllm_tensor_parallel_size, + pretrain=llm_cfg.model_name_or_path, + enable_prefix_caching=llm_cfg.enable_prefix_caching, + max_model_len=llm_cfg.prompt_max_len + llm_cfg.generate_max_len, + gpu_memory_utilization=llm_cfg.gpu_memory_utilization, + vllm_enable_sleep=llm_cfg.vllm_enable_sleep, + ) + + print(f'[Rank {rank}] Vllm engine successfully created!') + + from priorzero_datafactory import DataProcessor + data_processor = DataProcessor(rank=rank, + world_size=world_size, + vllm_engine=vllm_engine, + strategy=strategy, + model_path=llm_cfg.model_name_or_path, + exp_name=cfg.exp_name if rank == 0 else None, + ) + # 在collector中初始化data_processor 和prof对象 + collector.data_processor = data_processor + collector.prof = prof + evaluator.data_processor = data_processor + + policy_model = PolicyModel( + strategy=strategy, + pretrain=llm_cfg.model_name_or_path, + vllm_engine=vllm_engine, + max_steps=llm_cfg.max_steps + ) + from priorzero_trainer import PriorZeroLLMTrainer + trainer = PriorZeroLLMTrainer( + cfg=llm_cfg, + pretrain=llm_cfg.model_name_or_path, + strategy= strategy, + vllm_engine = vllm_engine, + policy_model=policy_model, + reference_model=ref_model, + exp_name=cfg.exp_name if rank == 0 else None, + tb_logger=tb_logger if rank == 0 else None, + llm_save_freq=llm_cfg.llm_save_freq + ) + + torch_dist_barrier_and_cuda_sync() + train_schedule = llm_cfg.train_schedule + train_alternate = train_schedule["alternate"] + current_phase = None + llm_collect_mode = None + if train_alternate: + current_phase = train_schedule["start_phase"] + last_wm_train_iter = 0 + last_llm_train_iter = 0 + llm_collect_mode = train_schedule["llm_collect_mode"] + + while True: + if collector.envstep >= max_env_step or learner.train_iter >= max_train_iter: + break + + # 1.评估阶段 + if learner.train_iter != 0 and evaluator.should_eval(wm_train_iter=learner.train_iter, llm_train_iter=policy_model.train_iter, phase=current_phase, env_step=collector.envstep): + logger.info(f"[Evaluator][Rank {rank}: Iter {learner.train_iter}] Evaluating...") + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.wake_up() + evaluator.eval(wm_train_iter=learner.train_iter, llm_train_iter=policy_model.train_iter, phase=current_phase, env_step=collector.envstep) + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.sleep() + + # 2.数据收集阶段 + if not train_alternate or (train_alternate and current_phase == "wm") or (train_alternate and current_phase == "llm" and llm_collect_mode != "no_collect"): + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.wake_up() + + new_data = collector.collect(train_iter=learner.train_iter, policy_kwargs={'temperature': 0.25, 'epsilon': 0.0}, phase=current_phase) + data_processor.get_llm_output_log(wm_train_iter=learner.train_iter, llm_train_iter=policy_model.train_iter) + + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.sleep() + + replay_buffer.push_game_segments(new_data) + replay_buffer.remove_oldest_data_to_fit() + + num_of_transitions = replay_buffer.get_num_of_transitions() + torch_dist_barrier_and_cuda_sync() + + # 3.world model训练阶段 + if llm_cfg.enable_world_model and (not train_alternate or (train_alternate and current_phase == "wm")): + if not (num_of_transitions > batch_size): + logger.warning(f'[WM Training] Data in replay_buffer is not sufficient: batch_size: {batch_size}, replay_buffer: {replay_buffer}. Continue to collect...') + cmd = 0 + else: + cmd = 1 + if min(all_gather_cmd(world_size=world_size, obj=cmd)) == 0: + continue + + update_per_collect = calculate_update_per_collect(cfg, new_data, world_size=world_size) + logger.info(f"[WM Training] Rank {rank} | Iter {learner.train_iter} | Updates: {update_per_collect}") + + for i in range(update_per_collect): + with prof.block("train_world_model", rank=rank): + train_data = replay_buffer.sample(batch_size, policy) + train_data.append(learner.train_iter) + + log_vars = learner.train(train_data, collector.envstep) + if cfg.policy.use_priority: + replay_buffer.update_priority(train_data, log_vars[0]['value_priority_orig']) + policy.recompute_pos_emb_diff_and_clear_cache() + if llm_cfg.enable_rft and train_alternate and learner.train_iter - last_wm_train_iter >= train_schedule["wm_update_iters"]: + current_phase = "llm" + last_wm_train_iter = learner.train_iter + if llm_collect_mode != "no_collect": + replay_buffer.mark_latest_transitions_consumed() + print(f"[WM Training][Rank {rank}] Switching to LLM training phase at wm iter: {learner.train_iter}") + continue + + # 4. llm 训练阶段 + if llm_cfg.enable_rft and (not train_alternate or (train_alternate and current_phase == "llm")): + new_num_of_transitions = replay_buffer.get_num_of_transitions() - replay_buffer.last_pos_in_transition + logger.info(f"[LLM Training] Rank {rank} | Total transitions: {num_of_transitions} | New transitions: {new_num_of_transitions}") + + if llm_collect_mode != "no_collect": + priorzero_batch = replay_buffer.fetch_latest_batch(batch_size=-1, policy=policy, select_last=True) + else: + priorzero_batch = replay_buffer.fetch_latest_batch(batch_size=256, policy=policy, select_last=False) + # 清理 policy的cahce,防止OOM + torch.cuda.empty_cache() + with prof.block("train_llm", rank=rank): + llm_need_sample_cnt = llm_cfg.train_batch_size * llm_cfg.max_rollout_staleness // world_size + flag, train_samples = data_processor.make_llm_train_samples(priorzero_batch, ddp=True, max_samples=llm_need_sample_cnt) + + if not flag: + local_llm_ready = 0 + else: + local_llm_ready = 1 + gathered_llm_ready = all_gather_cmd(world_size=world_size, obj=local_llm_ready) + + if min(gathered_llm_ready) == 0: + logger.info( + f"[Rank {rank}] Skip LLM training because not all ranks have enough samples. " + f"ready_flags={gathered_llm_ready}, local_ready={local_llm_ready}, required_samples_per_rank={llm_need_sample_cnt}, train_samples={len(train_samples[0])}" + ) + continue + + trainer.train_batch(train_samples, collect_env_steps=collector.envstep) + if llm_collect_mode != "no_collect": + replay_buffer.mark_latest_transitions_consumed() + + torch_dist_barrier_and_cuda_sync() + if llm_cfg.enable_world_model and train_alternate and trainer.global_step - last_llm_train_iter >= train_schedule["llm_update_iters"]: + current_phase = "wm" + last_llm_train_iter = trainer.global_step + data_processor.clear_statis() + print(f"[Rank {rank}] Switching to World Model training phase at llm iter: {trainer.global_step}") + +def main(): + """ + Main entry point with argument parsing. + """ + import argparse + + parser = argparse.ArgumentParser( + description='PriorZero Training with Auto Model Configuration', + formatter_class=argparse.RawDescriptionHelpFormatter, + epilog=""" +Examples: + # Use default model (qwen2.5-1.5b) + torchrun --nproc_per_node 2 priorzero_entry_sync.py + + # Use specific model + torchrun --nproc_per_node 2 priorzero_entry_sync.py --model qwen2.5-0.5b + torchrun --nproc_per_node 2 priorzero_entry_sync.py --model qwen2.5-7b + + # List all available models + python priorzero_entry_sync.py --list-models + + # Different environment + torchrun --nproc_per_node 2 priorzero_entry_sync.py --env_id zork1.z5 --model qwen2.5-1.5b + """ + ) + parser.add_argument('--env_id', type=str, default='detective.z5', help='Jericho game ID') + parser.add_argument('--seed', type=int, default=0, help='Random seed') + parser.add_argument('--max_iter', type=int, default=int(1e6), help='Max training iterations') + parser.add_argument('--quick_test', action='store_true', default=False, help='Use quick test config') + # Model selection + parser.add_argument('--model', type=str, default="qwen2.5-3b", choices=get_available_models()) + parser.add_argument('--enable_profile', action='store_true', default=False) + parser.add_argument('--use_cot', action='store_true', default=False) + args = parser.parse_args() + + model_key = args.model if args.model else "qwen2.5-1.5b" + print(f"\n{'='*80}") + print(f"PriorZero Training Configuration") + print(f"{'='*80}") + print(f"Environment: {args.env_id}") + print(f"Model: {model_key}") + print(f"Seed: {args.seed}") + print(f"Quick Test: {args.quick_test}") + print(f"use cot: {args.use_cot}") + print(f"enable_profile: {args.enable_profile}") + print(f"{'='*80}\n") + + if args.quick_test: + logger.info("Using quick test configuration") + main_cfg, create_cfg, llm_cfg = get_priorzero_debug_config( + args.env_id, args.seed, use_cot=args.use_cot, + exp_name=f'data_priorzero/priorzero_debug_{args.env_id}', + model_key=model_key, + ) + else: + main_cfg, create_cfg, llm_cfg = get_priorzero_config( + args.env_id, args.seed, use_cot=args.use_cot, + model_key=model_key, + multi_gpu=True + ) + + train_priorzero( + main_cfg, + create_cfg, + llm_cfg, + seed=args.seed, + max_train_iter=args.max_iter, + enable_profile=args.enable_profile, # 是否要对各个耗时部分进行 profile + ) + + +if __name__ == "__main__": + os.environ['TOKENIZERS_PARALLELISM'] = 'false' + main() diff --git a/zoo/jericho/priorzero/src/priorzero_evaluator.py b/zoo/jericho/priorzero/src/priorzero_evaluator.py new file mode 100644 index 000000000..f3e1322e5 --- /dev/null +++ b/zoo/jericho/priorzero/src/priorzero_evaluator.py @@ -0,0 +1,838 @@ +import copy +import json +import os +import time +from collections import namedtuple +from typing import Optional, Callable, Tuple, Dict, Any, List + +from collections import deque, defaultdict +import numpy as np +import torch +import torch.distributed as dist +import wandb +from ding.envs import BaseEnvManager +from ding.torch_utils import to_ndarray, to_item, to_tensor +from ding.utils import build_logger, EasyTimer +from ding.utils import get_world_size, get_rank, broadcast_object_list +from ding.worker.collector.base_serial_evaluator import ISerialEvaluator, VectorEvalMonitor +from easydict import EasyDict + +from lzero.mcts.buffer.game_segment import GameSegment +from lzero.mcts.utils import prepare_observation +import threading +from lzero.worker.muzero_evaluator import MuZeroEvaluator as OriginalEvaluator + + +def extract_raw_obs_text(obs_dict: Dict[str, Any]) -> str: + """Extract text observation from environment observation dictionary.""" + if 'raw_obs_text' in obs_dict: + return str(obs_dict['raw_obs_text']) + if 'observation_str' in obs_dict: + return str(obs_dict['observation_str']) + if 'observation' in obs_dict: + obs = obs_dict['observation'] + if isinstance(obs, str): + return obs + return str(obs_dict) + + +def extract_raw_obs_image(obs_dict: Dict[str, Any]) -> np.ndarray: + """Extract image observation from environment observation dictionary.""" + if 'observation' in obs_dict: + obs = obs_dict['observation'] + if isinstance(obs, np.ndarray): + return obs + raise ValueError(f"Cannot extract image from observation: {obs_dict.keys()}") + + +class PriorZeroEvaluator(OriginalEvaluator): + """ + PriorZero evaluator with three selectable eval modes: + 1) world_model: default UniZero eval + 2) world_model_llm_prior: inject llm_prior to MCTS root policy logits + 3) llm_prior_only: ignore world model and greedily pick best llm_prior action + """ + + def __init__(self, llm_config: Dict, data_processor=None, prior_generator=None, + obs_type: str = 'text', env_id: str = None, **kwargs) -> None: + super().__init__(**kwargs) + self.llm_cfg = llm_config + self.data_processor = data_processor + self.prior_generator = prior_generator + self.obs_type = obs_type + self.env_id = env_id or '' + + if self._rank == 0: + self._logger_eval_episode, _ = build_logger( + f'./{self._exp_name}/log/evaluator', "evaluator_episode_info", need_tb=False + ) + import logging + for handler in self._logger_eval_episode.handlers: + handler.setFormatter(logging.Formatter("%(message)s")) + + self.eval_mode = llm_config.eval_dict + self.eval_freq = self.eval_mode.eval_freq + self.wm_eval_freq = self.eval_mode.wm_eval_freq + self.llm_eval_freq = self.eval_mode.llm_eval_freq + self.llm_prior_temperature = llm_config.llm_prior_temperature + self.history_buffers = defaultdict( + lambda: deque(maxlen=self.llm_cfg.history_length) + ) + self._last_wm_eval_iter = 0 + self._last_llm_eval_iter = 0 + self._last_eval_envstep = 0 + + self._logger.info(f"[RANK {self._rank}] ✓ PriorZeroEvaluator initialized with vLLM engine") + self._logger.info(f"[RANK {self._rank}] - History length: {self.llm_cfg.history_length}") + + def should_eval(self, wm_train_iter: int, llm_train_iter, phase='wm', env_step: int = -1) -> bool: + """ + Determine whether it's time to run an evaluation. + + When ``env_step >= 0`` the decision is based on env-step frequency + (``wm_eval_freq_envsteps`` / ``llm_eval_freq_envsteps`` in eval_dict). + Otherwise falls back to the legacy iter-based logic for backward + compatibility. + """ + # --- New env-step-based trigger (preferred) --- + if env_step >= 0: + wm_freq_es = getattr(self.eval_mode, 'wm_eval_freq_envsteps', 0) + llm_freq_es = getattr(self.eval_mode, 'llm_eval_freq_envsteps', 0) + freq = wm_freq_es if (phase is None or phase == 'wm') else llm_freq_es + if freq > 0: + if env_step == self._last_eval_envstep: + return False + if (env_step - self._last_eval_envstep) < freq and env_step != 0: + return False + self._last_eval_envstep = env_step + return True + + # --- Legacy iter-based trigger (fallback) --- + if phase is None or phase == 'wm': + if wm_train_iter == self._last_wm_eval_iter: + return False + if (wm_train_iter - self._last_wm_eval_iter) < self.wm_eval_freq and wm_train_iter != 0: + return False + self._last_wm_eval_iter = wm_train_iter + return True + elif phase == 'llm': + if llm_train_iter == self._last_llm_eval_iter: + return False + if (llm_train_iter - self._last_llm_eval_iter) < self.llm_eval_freq and llm_train_iter != 0: + return False + self._last_llm_eval_iter = llm_train_iter + return True + else: + raise ValueError("") + + def _should_continue_eval(self, local_done: bool) -> bool: + """DDP-aware loop termination: continue while ANY rank still needs to work. + + With vLLM TP > 1 spanning DDP ranks, an early-exiting rank would leave its TP partner + deadlocked at a vllm collective. We all_reduce(MAX) a 0/1 flag so all ranks break together. + For TP=1 / single-process, falls back to local `not local_done`. + """ + tp_size = getattr(self.llm_cfg, 'vllm_tensor_parallel_size', 1) + if dist.is_initialized() and dist.get_world_size() > 1 and tp_size > 1: + flag = torch.tensor([0 if local_done else 1], dtype=torch.long, + device=torch.cuda.current_device()) + dist.all_reduce(flag, op=dist.ReduceOp.MAX) + return flag.item() > 0 + return not local_done + + def _save_eval_trajectories(self, completed_episodes: List[tuple], global_step: int, tag: str = "WM_LLMPrior") -> None: + """Save per-episode trajectory JSONs for post-hoc qualitative analysis. + + Each entry in completed_episodes is (level_id, total_reward, steps_list). + steps_list items: {obs, action, reward, mcts_info, info}. + Only called on rank 0. + """ + base_dir = os.path.join(f'./{self._exp_name}', 'eval_trajectories', f'step_{global_step}_{tag}') + level_counts: Dict[int, int] = {} + level_rewards: Dict[int, List[float]] = defaultdict(list) + + for level_id, total_reward, steps in completed_episodes: + lid = int(level_id) if level_id is not None else -1 + idx = level_counts.get(lid, 0) + level_counts[lid] = idx + 1 + level_rewards[lid].append(total_reward) + + level_dir = os.path.join(base_dir, f'level_{lid}') + os.makedirs(level_dir, exist_ok=True) + + traj = { + 'level_id': lid, + 'total_reward': total_reward, + 'episode_length': len(steps), + 'steps': [], + } + for s in steps: + step_record = { + 'obs': str(s.get('obs', ''))[:2000], + 'action': str(s.get('action', '')), + 'reward': float(s.get('reward', 0)), + } + info = s.get('info', {}) + if isinstance(info, dict): + step_record['data_idx'] = info.get('data_idx') + step_record['level_id'] = info.get('level_id') + + # --- Enriched fields: CoT, LLM prior, MCTS info --- + save_cot = getattr(self.eval_mode, 'save_llm_cot', True) + if save_cot: + # LLM CoT raw output + if s.get('llm_cot_raw') is not None: + step_record['llm_cot_raw'] = str(s['llm_cot_raw'])[:5000] + # LLM prompt + if s.get('llm_prompt') is not None: + step_record['llm_prompt'] = str(s['llm_prompt'])[:5000] + # LLM action probability distribution (normalized) + llm_probs = s.get('llm_action_probs') + if llm_probs: + step_record['llm_action_probs'] = { + str(a): float(p) for a, p in llm_probs.items() + } + # LLM policy (from eval_only_llm_prior path) + llm_policy = s.get('llm_policy') + if llm_policy: + step_record['llm_policy'] = { + str(a): float(p) for a, p in llm_policy.items() + } + # Valid actions list + va = s.get('valid_actions') + if va: + step_record['valid_actions'] = [str(a) for a in va] + # MCTS info (visit counts, prior distributions, etc.) + mcts = s.get('mcts_info') + if mcts and isinstance(mcts, dict): + mcts_serialized = {} + for key, value in mcts.items(): + if isinstance(value, dict): + mcts_serialized[str(key)] = { + str(a): float(v) if isinstance(v, (int, float)) else str(v) + for a, v in value.items() + } + else: + mcts_serialized[str(key)] = str(value) + step_record['mcts_info'] = mcts_serialized + + traj['steps'].append(step_record) + + with open(os.path.join(level_dir, f'traj_{idx}.json'), 'w') as f: + json.dump(traj, f, indent=2, ensure_ascii=False, default=str) + + index = { + 'global_step': global_step, + 'tag': tag, + 'n_episodes': len(completed_episodes), + 'levels': { + str(lid): {'n_traj': level_counts[lid], 'mean_reward': float(np.mean(level_rewards[lid]))} + for lid in sorted(level_rewards) + }, + } + with open(os.path.join(base_dir, 'index.json'), 'w') as f: + json.dump(index, f, indent=2, ensure_ascii=False) + self._logger.info(f"[EVALUATOR] Saved {len(completed_episodes)} trajectories to {base_dir}") + + def _log_per_level_tb(self, per_level_results: dict, tag_prefix: str, global_step: int) -> None: + """Log per-level rewards and summary to TensorBoard.""" + if not per_level_results or self._tb_logger is None: + return + all_level_means = [] + for level_id in sorted(per_level_results.keys()): + rewards = per_level_results[level_id] + mean_r = np.mean(rewards) + self._tb_logger.add_scalar(f'{tag_prefix}/level_{level_id}_reward', mean_r, global_step) + all_level_means.append(mean_r) + self._tb_logger.add_scalar(f'{tag_prefix}/level_mean', np.mean(all_level_means), global_step) + self._tb_logger.add_scalar(f'{tag_prefix}/level_std', np.std(all_level_means), global_step) + self._tb_logger.add_scalar(f'{tag_prefix}/level_min', np.min(all_level_means), global_step) + self._tb_logger.add_scalar(f'{tag_prefix}/level_max', np.max(all_level_means), global_step) + + def _log_agg_tb(self, info: dict, tag_prefix: str, global_step: int) -> None: + """Log aggregated eval metrics to TensorBoard.""" + if self._tb_logger is None: + return + for k in ['avg_envstep_per_episode', 'reward_mean', 'reward_std', 'reward_max', 'reward_min']: + if k in info: + self._tb_logger.add_scalar(f'{tag_prefix}/{k}', info[k], global_step) + + def eval(self, wm_train_iter: int = -1, llm_train_iter: int = -1, phase: str = "wm", env_step: int = -1) -> Tuple[bool, Dict[str, Any]]: + modes = [] + wm_per_level = {} + wm_llm_per_level = {} + llm_per_level = {} + + # Mode 1: Pure WM+MCTS — now runs in ALL phases (no phase guard) + if self.eval_mode.world_model: + world_model_info, wm_per_level = self.eval_wm_only() + modes.append(("WM", world_model_info)) + tp_size = getattr(self.llm_cfg, 'vllm_tensor_parallel_size', 1) + if dist.is_initialized() and dist.get_world_size() > 1 and tp_size > 1: + dist.barrier() + wm_llm_completed_episodes = [] + if self.eval_mode.world_model_llm_prior: + world_model_llm_prior_info, wm_llm_eval_episode_info, wm_llm_per_level, wm_llm_completed_episodes = self.eval_with_llm_prior() + modes.append(("WM_LLMPrior", world_model_llm_prior_info)) + + if self.eval_mode.llm_prior: + llm_prior_info, llm_eval_episode_info, llm_per_level = self.eval_only_llm_prior() + modes.append(("LLMPrior", llm_prior_info)) + + if self._rank != 0: + return + + # --- Save evaluation trajectories for post-hoc analysis --- + step_val = wm_train_iter if (phase == 'wm' or phase is None) else llm_train_iter + if wm_llm_completed_episodes: + self._save_eval_trajectories(wm_llm_completed_episodes, step_val, tag="WM_LLMPrior") + + # --- Episode-level text logging --- + if self.eval_mode.world_model_llm_prior and wm_llm_eval_episode_info and len(wm_llm_eval_episode_info[0]) > 0: + self._logger_eval_episode.info("="*100) + self._logger_eval_episode.info("="*10 + f"[WM_LLM] | episode_avg_steps={len(wm_llm_eval_episode_info[0])} | episode_return={wm_llm_eval_episode_info[0][-1]['info']['score'].item()} " + "="*10) + for step, info in enumerate(wm_llm_eval_episode_info[0]): + obs, action, reward, mcts_info = info['obs'].replace("\n",""), info['action'], info['reward'], info['mcts_info'] + self._logger_eval_episode.info(f"[Step {step:03d}] obs: {obs}") + self._logger_eval_episode.info(f'action="{action}" | reward={reward}') + self._logger_eval_episode.info("MCTS:") + for key, value in mcts_info.items(): + items = list(value.items()) + action_str = " | ".join( + f"{a}({v:.3f})" if isinstance(v, float) else f"{a}({v})" + for a, v in items + ) + self._logger_eval_episode.info(f" {key}:") + self._logger_eval_episode.info(f" {action_str}") + self._logger_eval_episode.info("-" * 100) + self._logger_eval_episode.info("="*100) + + if llm_eval_episode_info is not None and self.obs_type == 'text': + self._logger_eval_episode.info("="*100) + self._logger_eval_episode.info("="*10 + f"[LLM] | episode_avg_steps={len(llm_eval_episode_info[0])} | episode_return={llm_eval_episode_info[0][-1]['info']['score'].item()} " + "="*10) + for step, info in enumerate(llm_eval_episode_info[0]): + obs, action, reward, llm_policy = info['obs'].replace("\n",""), info['action'], info['reward'], info['llm_policy'] + self._logger_eval_episode.info(f"[Step {step:03d}] obs: {obs}") + self._logger_eval_episode.info(f'action="{action}" | reward={reward}') + items = list(llm_policy.items()) + action_str = " | ".join( + f"{a}({v:.3f})" if isinstance(v, float) else f"{a}({v})" + for a, v in items + ) + self._logger_eval_episode.info("llm_policy:") + self._logger_eval_episode.info(f" {action_str}") + self._logger_eval_episode.info("-" * 100) + self._logger_eval_episode.info("="*100) + + # Image mode: structured summary log + if self.obs_type == 'image': + self._logger_eval_episode.info("=" * 80) + self._logger_eval_episode.info(f"[Eval Summary] obs_type=image | env={self.env_id}") + self._logger_eval_episode.info("-" * 80) + for tag, info in modes: + self._logger_eval_episode.info( + f" [{tag}] reward_mean={info.get('reward_mean', 0):.2f} | " + f"reward_max={info.get('reward_max', 0):.2f} | " + f"reward_min={info.get('reward_min', 0):.2f} | " + f"avg_steps={info.get('avg_envstep_per_episode', 0):.1f}" + ) + if wm_llm_eval_episode_info is not None and len(wm_llm_eval_episode_info[0]) > 0: + ep = wm_llm_eval_episode_info[0] + ep_return = ep[-1]['info'].get('eval_episode_return', ep[-1]['info'].get('score', 'N/A')) + self._logger_eval_episode.info(f" [WM_VLPrior ep0] steps={len(ep)} | return={ep_return}") + if llm_eval_episode_info is not None and len(llm_eval_episode_info[0]) > 0: + ep = llm_eval_episode_info[0] + ep_return = ep[-1]['info'].get('eval_episode_return', ep[-1]['info'].get('score', 'N/A')) + self._logger_eval_episode.info(f" [VLPrior ep0] steps={len(ep)} | return={ep_return}") + self._logger_eval_episode.info("=" * 80) + + keys = ['avg_envstep_per_episode', 'reward_mean', 'reward_std', 'reward_max', 'reward_min'] + for k in keys: + if world_model_info is not None: + self._tb_logger.add_scalar(f'{self._instance_name}_iter/{k}_WM', world_model_info[k], train_iter) + self._tb_logger.add_scalar(f'{self._instance_name}_step/{k}_WM', world_model_info[k], envstep) + if world_model_llm_prior_info is not None: + self._tb_logger.add_scalar(f'{self._instance_name}_iter/{k}_WM_LLMPrior', world_model_llm_prior_info[k], train_iter) + self._tb_logger.add_scalar(f'{self._instance_name}_step/{k}_WM_LLMPrior', world_model_llm_prior_info[k], envstep) + if llm_prior_info is not None: + self._tb_logger.add_scalar(f'{self._instance_name}_iter/{k}_LLMPrior', llm_prior_info[k], train_iter) + self._tb_logger.add_scalar(f'{self._instance_name}_step/{k}_LLMPrior', llm_prior_info[k], envstep) + + return stop_flag, best_reward + + # ================================================================== + # eval_with_llm_prior: WM + VL/LLM prior → MCTS + # ================================================================== + + def eval_with_llm_prior(self) -> Dict[str, Any]: + n_episode = self._default_n_episode + assert n_episode is not None, "Please specify the number of evaluation episodes (n_episode)." + envstep_count = 0 + completed_episodes: List[tuple] = [] + # Hard counter independent of VectorEvalMonitor's per-env deque-fullness check; guards + # against eval hanging when episodes are unevenly distributed across envs. + total_finishes = 0 + eval_monitor = VectorEvalMonitor(self._env.env_num, n_episode) + env_nums = self._env.env_num + + eval_episode_info = [[] for _ in range(env_nums)] + # aligned with ScalingInter-RL: track per-level results for TensorBoard + per_level_results = defaultdict(list) + + self._env.reset() + self.history_buffers.clear() + self._policy.reset(task_id=self.task_id) + + init_obs = self._env.ready_obs + + retry_waiting_time = 0.001 + while len(init_obs.keys()) != self._env_num: + self._logger.info(f"[RANK {self._rank}] Waiting for all environments to reset. Current ready envs: {list(init_obs.keys())}") + time.sleep(retry_waiting_time) + init_obs = self._env.ready_obs + + action_mask_dict = {i: to_ndarray(init_obs[i]['action_mask']) for i in range(env_nums)} + to_play_dict = {i: to_ndarray(init_obs[i]['to_play']) for i in range(env_nums)} + + timestep_dict = {} + for i in range(env_nums): + if 'timestep' not in init_obs[i]: + self._logger.debug(f"'timestep' missing in init_obs[{i}], using -1") + timestep_dict[i] = to_ndarray(init_obs[i].get('timestep', -1)) + + dones = np.array([False for _ in range(env_nums)]) + + game_segments = [ + GameSegment( + self._env.action_space, + game_segment_length=self.policy_config.game_segment_length, + config=self.policy_config, + task_id=self.task_id + ) for _ in range(env_nums) + ] + for i in range(env_nums): + game_segments[i].reset( + [to_ndarray(init_obs[i]['observation']) for _ in range(self.policy_config.model.frame_stack_num)] + ) + + ready_env_id = set() + remain_episode = n_episode + eps_steps_lst = np.zeros(env_nums) + with self._timer: + while True: + local_done = (total_finishes >= n_episode) or eval_monitor.is_finished() + if not self._should_continue_eval(local_done): + break + if local_done: + # Drain mode: this rank already collected n_episode results, but must keep + # issuing matched vllm.generate calls so its TP partners can finish theirs. + self.data_processor.drain_vllm_iter() + continue + # Check if a timeout has occurred. + if self.stop_event.is_set(): + self._logger.info("[RANK {self._rank}] [EVALUATOR]: Evaluation aborted due to timeout.") + break + + obs = self._env.ready_obs + new_available_env_id = set(obs.keys()).difference(ready_env_id) + ready_env_id = ready_env_id.union(set(list(new_available_env_id)[:remain_episode])) + remain_episode -= min(len(new_available_env_id), remain_episode) + + if not ready_env_id: + continue + + # Prepare stacked observations and other inputs for the policy. + stack_obs = {env_id: game_segments[env_id].get_obs() for env_id in ready_env_id} + stack_obs = list(stack_obs.values()) + action_mask = [action_mask_dict[env_id] for env_id in ready_env_id] + to_play = [to_play_dict[env_id] for env_id in ready_env_id] + timestep = [timestep_dict[env_id] for env_id in ready_env_id] + + stack_obs = to_ndarray(stack_obs) + stack_obs = prepare_observation(stack_obs, self.policy_config.model.model_type) + stack_obs = torch.from_numpy(stack_obs).to(self.policy_config.device).float() + + # ============================================ + # Get VL/LLM Prior + # ============================================ + raw_obs_list = [] + histories_list = [] + valid_actions_list = [] + for env_id in sorted(list(ready_env_id)): + raw_obs_text = obs[env_id]['raw_obs_text'] + raw_obs_list.append(raw_obs_text) + + history = list(self.history_buffers[env_id]) + histories_list.append(history) + + valid_actions = obs[env_id].get('valid_actions', []) + valid_actions_list.append(valid_actions) + + llm_prior_per_seq, llm_prior_per_tok, prefix_cots, full_cot_outputs = self.data_processor.get_llm_prior( + states=raw_obs_list, + valid_actions_list=valid_actions_list, # [PRIORZERO] Pass valid actions + histories=histories_list, + return_cot=True # Request CoT prefixes for reuse in training + ) + + # Build per-env lookup for CoT/prompt data (aligned with sorted ready_env_id) + sorted_ready = sorted(list(ready_env_id)) + _llm_cot_by_env = {} + _llm_prompt_by_env = {} + _llm_prior_raw_by_env = {} # unscaled log-probs + _valid_actions_by_env = {} + for idx, env_id in enumerate(sorted_ready): + _valid_actions_by_env[env_id] = valid_actions_list[idx] + _llm_prior_raw_by_env[env_id] = dict(llm_prior_per_seq[idx]) # copy before scaling + if full_cot_outputs and idx < len(full_cot_outputs): + _llm_cot_by_env[env_id] = full_cot_outputs[idx] + else: + _llm_cot_by_env[env_id] = None + if llm_prior_per_tok and idx < len(llm_prior_per_tok): + _llm_prompt_by_env[env_id] = llm_prior_per_tok[idx].get('prompt', None) + else: + _llm_prompt_by_env[env_id] = None + + for env_id, llm_prior in enumerate(llm_prior_per_seq): + scaled_llm_prior = self.apply_temperature_scaling(llm_prior, return_logprobs=True) + llm_prior_per_seq[idx] = scaled_llm_prior + + policy_kwargs_forward = { + 'llm_prior_logprob': llm_prior_per_seq, + 'valid_actions_list': valid_actions_list, + } + if self.task_id is not None: + policy_kwargs_forward['task_id'] = self.task_id + + # ============================================================== + # Policy Forward Pass + # ============================================================== + policy_output, mcts_info = self._policy.forward( + data=stack_obs, action_mask=action_mask, + to_play=to_play, ready_env_id=ready_env_id, + timestep=timestep, **policy_kwargs_forward + ) + actions_with_env_id = {k: v['action'] for k, v in policy_output.items()} + distributions_dict_with_env_id = {k: v['visit_count_distributions'] for k, v in policy_output.items()} + value_dict_with_env_id = {k: v['searched_value'] for k, v in policy_output.items()} + pred_value_dict_with_env_id = {k: v['predicted_value'] for k, v in policy_output.items()} + timestep_dict_with_env_id = {k: v.get('timestep', -1) for k, v in policy_output.items()} + visit_entropy_dict_with_env_id = {k: v['visit_count_distribution_entropy'] for k, v in policy_output.items()} + + actions, distributions_dict, value_dict, pred_value_dict = {}, {}, {}, {} + visit_entropy_dict = {} + for index, env_id in enumerate(ready_env_id): + actions[env_id] = actions_with_env_id.pop(env_id) + distributions_dict[env_id] = distributions_dict_with_env_id.pop(env_id) + value_dict[env_id] = value_dict_with_env_id.pop(env_id) + pred_value_dict[env_id] = pred_value_dict_with_env_id.pop(env_id) + timestep_dict[env_id] = timestep_dict_with_env_id.pop(env_id) + visit_entropy_dict[env_id] = visit_entropy_dict_with_env_id.pop(env_id) + + # ============================================================== + # Environment Interaction + # ============================================================== + timesteps = self._env.step(actions) + timesteps = to_tensor(timesteps, dtype=torch.float32) + for env_id, episode_timestep in timesteps.items(): + obs_new, reward, done, info = episode_timestep.obs, episode_timestep.reward, episode_timestep.done, episode_timestep.info + + action_str = self._action_index_to_str(actions[env_id], valid_actions_list, info) + obs_repr = self._extract_obs(obs[env_id]) if self.obs_type == 'text' else f"image_{env_id}" + eval_episode_info[env_id].append({ + "obs": obs_repr, + "action": action_str, + "reward": float(reward), + "mcts_info": mcts_info[env_id], + "info": info, + # --- enriched fields for trajectory analysis --- + "llm_cot_raw": _llm_cot_by_env.get(env_id), + "llm_prompt": _llm_prompt_by_env.get(env_id), + "llm_action_probs": _llm_prior_raw_by_env.get(env_id, {}), + "valid_actions": _valid_actions_by_env.get(env_id, []), + }) + # Update history with absolute timestep + raw_obs_for_history = self._extract_obs(obs[env_id]) + abs_timestep = int(timestep_dict[env_id]) if int(timestep_dict[env_id]) >= 0 else int(eps_steps_lst[env_id]) + self.history_buffers[env_id].append((raw_obs_for_history, action_str, float(reward), abs_timestep)) + + eps_steps_lst[env_id] += 1 + if self._policy.get_attribute('cfg').type in ['unizero', 'sampled_unizero', 'priorzero']: + self._policy.reset(env_id=env_id, current_steps=eps_steps_lst[env_id], reset_init_data=False) + + game_segments[env_id].append( + actions[env_id], to_ndarray(obs_new['observation']), reward, action_mask_dict[env_id], + to_play_dict[env_id], timestep_dict[env_id] + ) + + action_mask_dict[env_id] = to_ndarray(obs_new['action_mask']) + to_play_dict[env_id] = to_ndarray(obs_new['to_play']) + timestep_dict[env_id] = to_ndarray(obs_new.get('timestep', -1)) + + dones[env_id] = done + if episode_timestep.done: + self._policy.reset([env_id]) + reward = episode_timestep.info.get('score', episode_timestep.info.get('eval_episode_return', 0)) + saved_info = {'eval_episode_return': reward} + if 'episode_info' in episode_timestep.info: + saved_info.update(episode_timestep.info['episode_info']) + # Only count up to n_episode; drain-mode iters never reach here (body skipped). + if total_finishes < n_episode: + eval_monitor.update_info(env_id, saved_info) + eval_monitor.update_reward(env_id, reward) + total_finishes += 1 + + # aligned with ScalingInter-RL: record per-level result + level_id = episode_timestep.info.get('level_id', None) + if level_id is not None: + per_level_results[int(level_id)].append(float(reward)) + + completed_episodes.append((level_id, float(reward), list(eval_episode_info[env_id]))) + eval_episode_info[env_id] = [] + + # Remove BEFORE the inner refill: only then does + # `init_obs.keys() - ready_env_id` actually include this env_id. + ready_env_id.remove(env_id) + + if n_episode > self._env_num: + init_obs = self._env.ready_obs + while len(init_obs.keys()) != self._env_num: + self._logger.info(f"Waiting for env {env_id} to reset. Current ready envs: {list(init_obs.keys())}") + time.sleep(retry_waiting_time) + init_obs = self._env.ready_obs + + new_available_env_id = set(init_obs.keys()).difference(ready_env_id) + ready_env_id = ready_env_id.union(set(list(new_available_env_id)[:remain_episode])) + remain_episode -= min(len(new_available_env_id), remain_episode) + + action_mask_dict[env_id] = to_ndarray(init_obs[env_id]['action_mask']) + to_play_dict[env_id] = to_ndarray(init_obs[env_id]['to_play']) + timestep_dict[env_id] = to_ndarray(init_obs[env_id].get('timestep', -1)) + + game_segments[env_id] = GameSegment( + self._env.action_space, + game_segment_length=self.policy_config.game_segment_length, + config=self.policy_config, + task_id=self.task_id + ) + game_segments[env_id].reset( + [init_obs[env_id]['observation'] for _ in range(self.policy_config.model.frame_stack_num)] + ) + + eps_steps_lst[env_id] = 0 + self._policy.reset([env_id]) + + envstep_count += 1 + + duration = self._timer.value + episode_return = eval_monitor.get_episode_return() + info = { + 'avg_envstep_per_episode': envstep_count / n_episode if n_episode > 0 else 0, + 'reward_mean': np.mean(episode_return), + 'reward_std': np.std(episode_return), + 'reward_max': np.max(episode_return), + 'reward_min': np.min(episode_return), + } + return info, eval_episode_info, dict(per_level_results), completed_episodes + + def eval_only_llm_prior(self) -> Dict[str, Any]: + n_episode = self._default_n_episode + assert n_episode is not None, "Please specify the number of evaluation episodes (n_episode)." + envstep_count = 0 + total_finishes = 0 + env_nums = self._env.env_num + + eval_episode_info = [[] for _ in range(env_nums)] + per_level_results = defaultdict(list) + + self._env.reset() + self.history_buffers.clear() + + dones = np.array([False for _ in range(env_nums)]) + ready_env_id = set(range(env_nums)) + remain_episode = n_episode + episode_return = [] + + retry_waiting_time = 0.001 + + init_obs = self._env.ready_obs + while len(init_obs.keys()) != self._env_num: + time.sleep(retry_waiting_time) + init_obs = self._env.ready_obs + + while True: + local_done = (total_finishes >= n_episode) + if not self._should_continue_eval(local_done): + break + if local_done: + self.data_processor.drain_vllm_iter() + continue + + obs = self._env.ready_obs + # ============================================ + # Get VL/LLM Prior + # ============================================ + raw_obs_list = [] + histories_list = [] + valid_actions_list = [] + for env_id in sorted(list(ready_env_id)): + raw_obs_text = obs[env_id]['raw_obs_text'] + raw_obs_list.append(raw_obs_text) + + history = list(self.history_buffers[env_id]) + histories_list.append(history) + + valid_actions = obs[env_id].get('valid_actions', []) + valid_actions_list.append(valid_actions) + + llm_prior_per_seq, llm_prior_per_tok, prefix_cots, full_cot_outputs = self.data_processor.get_llm_prior( + states=raw_obs_list, + valid_actions_list=valid_actions_list, + histories=histories_list, + return_cot=True + ) + # Build per-env lookup for CoT/prompt data + sorted_ready_llm = sorted(list(ready_env_id)) + llm_cot_by_env = {} + llm_prompt_by_env = {} + for idx, env_id in enumerate(sorted_ready_llm): + if full_cot_outputs is not None and idx < len(full_cot_outputs): + llm_cot_by_env[env_id] = full_cot_outputs[idx] + else: + llm_cot_by_env[env_id] = None + if llm_prior_per_tok is not None and idx < len(llm_prior_per_tok): + llm_prompt_by_env[env_id] = llm_prior_per_tok[idx].get('prompt', None) + else: + llm_prompt_by_env[env_id] = None + + actions = {env_id: None for env_id in sorted(list(ready_env_id))} + llm_policy = {env_id: {} for env_id in sorted(list(ready_env_id))} + + for env_id, llm_prior, valid_actions in zip(sorted(list(ready_env_id)), llm_prior_per_seq, valid_actions_list): + # llm_prior can be a dict (text) or np.ndarray (image) + if isinstance(llm_prior, np.ndarray): + # Image mode: prior is an array of probs, pick argmax + actions[env_id] = int(np.argmax(llm_prior)) + for i, action_name in enumerate(valid_actions): + llm_policy[env_id][action_name] = float(llm_prior[i]) if i < len(llm_prior) else 0.0 + elif isinstance(llm_prior, dict): + # Text mode: prior is a dict of action_str -> logprob + if len(llm_prior) == 1: + assert len(valid_actions) == 0 + actions[env_id] = 0 + continue + if 'go' in llm_prior and 'go' not in valid_actions: + llm_prior.pop('go') + action_str_select, max_logprob = "", float(-1e9) + for action_str, logprob in llm_prior.items(): + llm_policy[env_id][action_str] = np.exp(logprob) + if logprob > max_logprob: + action_str_select = action_str + max_logprob = logprob + all_values = [v for _, v in llm_policy[env_id].items()] + for k, _ in llm_policy[env_id].items(): + llm_policy[env_id][k] /= sum(all_values) + actions[env_id] = valid_actions.index(action_str_select) + else: + # Fallback: uniform random + actions[env_id] = 0 + + # ============================================ + timesteps = self._env.step(actions) + timesteps = to_tensor(timesteps, dtype=torch.float32) + for env_id, episode_timestep in timesteps.items(): + obs_new, reward, done, info = episode_timestep.obs, episode_timestep.reward, episode_timestep.done, episode_timestep.info + + action_str = self._action_index_to_str(actions[env_id], valid_actions_list, info) + obs_repr = self._extract_obs(obs[env_id]) if self.obs_type == 'text' else f"image_{env_id}" + eval_episode_info[env_id].append({ + "obs": obs_repr, + "action": action_str, + "reward": float(reward), + "llm_policy": llm_policy[env_id], + "info": info, + # --- CoT / LLM prior enrichment --- + "llm_cot_raw": llm_cot_by_env.get(env_id), + "llm_prompt": llm_prompt_by_env.get(env_id), + "valid_actions": valid_actions_list[sorted_ready_llm.index(env_id)] if env_id in sorted_ready_llm else [], + }) + raw_obs_for_history = self._extract_obs(obs[env_id]) + self.history_buffers[env_id].append((raw_obs_for_history, action_str, float(reward), int(eps_steps_lst[env_id]))) + + eps_steps_lst[env_id] += 1 + dones[env_id] = done + if episode_timestep.done: + ready_env_id.discard(env_id) + if total_finishes < n_episode: + episode_return.append(info['score']) + total_finishes += 1 + + level_id = info.get('level_id', None) + if level_id is not None: + per_level_results[int(level_id)].append(float(info['score'])) + + if n_episode > self._env_num and total_finishes < n_episode: + init_obs = self._env.ready_obs + while len(init_obs.keys()) != self._env_num: + time.sleep(retry_waiting_time) + init_obs = self._env.ready_obs + + new_available_env_id = set(init_obs.keys()).difference(ready_env_id) + ready_env_id = ready_env_id.union(set(list(new_available_env_id)[:remain_episode])) + remain_episode -= min(len(new_available_env_id), remain_episode) + + self.history_buffers[env_id].clear() + dones[env_id] = False + eval_episode_info[env_id] = [] + + envstep_count += 1 + info = { + 'avg_envstep_per_episode': envstep_count / n_episode if n_episode > 0 else 0, + 'reward_mean': np.mean(episode_return), + 'reward_std': np.std(episode_return), + 'reward_max': np.max(episode_return), + 'reward_min': np.min(episode_return), + } + return info, eval_episode_info, dict(per_level_results) + + def apply_temperature_scaling(self, logprobs_dict: dict, return_logprobs: bool = True) -> dict: + """ + Apply temperature scaling. Handles both dict (text) and ndarray (image) formats. + """ + import math + T = self.llm_prior_temperature + + # Image mode: ndarray of probs → convert to log-probs, scale, convert back + if isinstance(logprobs_input, np.ndarray): + log_probs = np.log(logprobs_input + 1e-10) + if T <= 1e-8: + result = np.zeros_like(log_probs) + result[np.argmax(log_probs)] = 0.0 # log(1)=0 + result[result == 0] = -1e10 + result[np.argmax(logprobs_input)] = 0.0 + return result if return_logprobs else np.exp(result) + scaled = log_probs / T + scaled -= scaled.max() + log_sum_exp = np.log(np.sum(np.exp(scaled))) + normalized = scaled - log_sum_exp + return normalized if return_logprobs else np.exp(normalized) + + # Text mode: dict of action_str -> logprob + if isinstance(logprobs_input, dict): + if T <= 1e-8: + max_key = max(logprobs_input, key=logprobs_input.get) + return {k: (0.0 if k != max_key else 1.0) for k in logprobs_input} + + scaled_logits = {k: v / T for k, v in logprobs_input.items()} + max_val = max(scaled_logits.values()) + sum_exp = sum(math.exp(v - max_val) for v in scaled_logits.values()) + log_sum_exp = math.log(sum_exp) + max_val + + result = {} + for k, v in scaled_logits.items(): + normalized_logprob = v - log_sum_exp + result[k] = normalized_logprob if return_logprobs else math.exp(normalized_logprob) + return result + + # Fallback + return logprobs_input diff --git a/zoo/jericho/priorzero/src/priorzero_policy.py b/zoo/jericho/priorzero/src/priorzero_policy.py new file mode 100644 index 000000000..795113dca --- /dev/null +++ b/zoo/jericho/priorzero/src/priorzero_policy.py @@ -0,0 +1,569 @@ +import asyncio +import copy +import inspect +import re +import sys +import logging +from pathlib import Path +from typing import List, Dict, Any, Tuple, Union, Optional +from collections import defaultdict + +import numpy as np +import torch +import torch.distributed as dist +import torch.nn.functional as F +from ding.utils import POLICY_REGISTRY, allreduce +from ding.model import model_wrap +import os + +# Import from local LightZero +from lzero.policy.unizero import UniZeroPolicy as OriginalUniZeroPolicy +from lzero.policy import phi_transform, InverseScalarTransform, scalar_transform, DiscreteSupport +from lzero.policy import to_torch_float_tensor,mz_network_output_unpack, prepare_obs +from lzero.policy.utils import select_action +from lzero.mcts import UniZeroMCTSCtree as MCTSCtree +from lzero.entry.utils import initialize_zeros_batch +import lzero.model.unizero_model + +@POLICY_REGISTRY.register('priorzero', force_overwrite=True) +class PriorZeroPolicy(OriginalUniZeroPolicy): + def __init__(self, cfg: Dict, model: torch.nn.Module = None, enable_field: List[str] = None, **kwargs): + super().__init__(cfg, model, enable_field) + self.llm_cfg = kwargs.get('llm_cfg', None) + + def _init_learn(self) -> None: + super()._init_learn() + logging.info("✓ UniZero World Model and optimizer initialized") + + def _forward_learn(self, data: Tuple[torch.Tensor]) -> Dict[str, Union[float, int]]: + self._learn_model.train() + self._target_model.train() + + current_batch, target_batch, train_iter = data + + # CoT reuse optimization: unpack cot_prefix_list (12 elements total) + obs_batch_ori, action_batch, target_action_batch, mask_batch, batch_index_tensor, weights, make_time, timestep_batch, raw_obs_list, history_obs_list, llm_prior_per_tok_list, cot_prefix_list, llm_action_list = current_batch + target_reward, target_value, target_policy = target_batch + + obs_batch, obs_target_batch = prepare_obs(obs_batch_ori, self._cfg) + action_batch = torch.from_numpy(action_batch).to(self._cfg.device).unsqueeze( + -1).long() + timestep_batch = torch.from_numpy(timestep_batch).to(self._cfg.device).unsqueeze( + -1).long() + + data_list = [mask_batch, target_reward, target_value, target_policy, weights] + (mask_batch, target_reward, target_value, target_policy, weights) = to_torch_float_tensor(data_list, self._cfg.device) + + batch_size = self._cfg.batch_size + target_reward = target_reward.view(batch_size, -1) + target_value = target_value.view(batch_size, -1) + + transformed_target_reward = scalar_transform(target_reward) + transformed_target_value = scalar_transform(target_value) + + # Convert to categorical distribution (for distributional RL) + target_reward_categorical = phi_transform(self.reward_support, transformed_target_reward) + target_value_categorical = phi_transform(self.value_support, transformed_target_value) + + batch_for_gpt = { + 'actions': action_batch.squeeze(-1), + 'timestep': timestep_batch.squeeze(-1), + 'rewards': target_reward_categorical[:, :-1], + 'target_value': target_value_categorical[:, :-1], + 'target_policy': target_policy[:, :-1], + } + if isinstance(self._cfg.model.observation_shape, int) or len(self._cfg.model.observation_shape) == 1: + batch_for_gpt['observations'] = torch.cat((obs_batch, obs_target_batch), dim=1).reshape( + self._cfg.batch_size, -1, self._cfg.model.observation_shape) + elif len(self._cfg.model.observation_shape) == 3: + batch_for_gpt['observations'] = torch.cat((obs_batch, obs_target_batch), dim=1).reshape( + self._cfg.batch_size, -1, *self._cfg.model.observation_shape) + + batch_for_gpt['mask_padding'] = mask_batch == 1.0 + batch_for_gpt['observations'] = batch_for_gpt['observations'][:, :-1] + batch_for_gpt['mask_padding'] = batch_for_gpt['mask_padding'][:, :-1] + batch_for_gpt['ends'] = torch.zeros(batch_for_gpt['mask_padding'].shape, dtype=torch.long, device=self._cfg.device) + batch_for_gpt['scalar_target_value'] = target_value + + wm_losses, pred_values = self._learn_model.world_model.compute_loss( + batch_for_gpt, + self._target_model.world_model.tokenizer, + self.value_inverse_scalar_transform_handle, + ) + + wm_total_loss = (weights * wm_losses.loss_total).mean() + + self._optimizer_world_model.zero_grad() + wm_total_loss.backward() + wm_grad_norm = torch.nn.utils.clip_grad_norm_( + self._learn_model.world_model.parameters(), + self._cfg.grad_clip_value + ) + if self._cfg.multi_gpu: + # Only sync world_model gradients (other params have None grad) + for p in self._learn_model.world_model.parameters(): + if p.grad is not None: + allreduce(p.grad.data) + self._optimizer_world_model.step() + self._target_model.update(self._learn_model.state_dict()) + + intermediate_losses = wm_losses.intermediate_losses + obs_loss = intermediate_losses.get('loss_obs', torch.tensor(0.0)) + reward_loss = intermediate_losses.get('loss_rewards', torch.tensor(0.0)) + policy_loss = intermediate_losses.get('loss_policy', torch.tensor(0.0)) + value_loss = intermediate_losses.get('loss_value', torch.tensor(0.0)) + latent_recon_loss = intermediate_losses.get('latent_recon_loss', torch.tensor(0.0)) + perceptual_loss = intermediate_losses.get('perceptual_loss', torch.tensor(0.0)) + orig_policy_loss = intermediate_losses.get('orig_policy_loss', torch.tensor(0.0)) + policy_entropy = intermediate_losses.get('policy_entropy', torch.tensor(0.0)) + first_step_losses = intermediate_losses.get('first_step_losses', {}) + middle_step_losses = intermediate_losses.get('middle_step_losses', {}) + last_step_losses = intermediate_losses.get('last_step_losses', {}) + + latent_state_l2_norms = intermediate_losses.get('latent_state_l2_norms', torch.tensor(0.0)) + latent_action_l2_norms = intermediate_losses.get('latent_action_l2_norms', 0.0) + + # Logits statistics + logits_value_mean = intermediate_losses.get('logits_value_mean', 0.0) + logits_value_max = intermediate_losses.get('logits_value_max', 0.0) + logits_value_min = intermediate_losses.get('logits_value_min', 0.0) + logits_policy_mean = intermediate_losses.get('logits_policy_mean', 0.0) + logits_policy_max = intermediate_losses.get('logits_policy_max', 0.0) + logits_policy_min = intermediate_losses.get('logits_policy_min', 0.0) + + # Temperature parameters + temperature_value = intermediate_losses.get('temperature_value', 0.0) + temperature_reward = intermediate_losses.get('temperature_reward', 0.0) + temperature_policy = intermediate_losses.get('temperature_policy', 0.0) + + # Value priority for prioritized replay + value_priority_tensor = intermediate_losses.get('value_priority', torch.tensor([0.0])) + value_priority_np = value_priority_tensor.detach().cpu().numpy() + 1e-6 + + # Compute target policy entropy (for analysis) + valid_target_policy = batch_for_gpt['target_policy'][batch_for_gpt['mask_padding']] + target_policy_entropy = -torch.sum(valid_target_policy * torch.log(valid_target_policy + 1e-9), dim=-1) + average_target_policy_entropy = target_policy_entropy.mean() + + # Build comprehensive log dict (aligned with UniZero) + log_dict = { + # ============ Core Losses ============ + 'wm_total_loss': wm_total_loss.item(), + 'wm_obs_loss': obs_loss.item() if torch.is_tensor(obs_loss) else obs_loss, + 'wm_reward_loss': reward_loss.item() if torch.is_tensor(reward_loss) else reward_loss, + 'wm_policy_loss': policy_loss.item() if torch.is_tensor(policy_loss) else policy_loss, + 'wm_value_loss': value_loss.item() if torch.is_tensor(value_loss) else value_loss, + 'wm_latent_recon_loss': latent_recon_loss.item() if torch.is_tensor(latent_recon_loss) else latent_recon_loss, + 'wm_perceptual_loss': perceptual_loss.item() if torch.is_tensor(perceptual_loss) else perceptual_loss, + 'wm_orig_policy_loss': orig_policy_loss.item() if torch.is_tensor(orig_policy_loss) else orig_policy_loss, + 'wm_policy_entropy': policy_entropy.item() if torch.is_tensor(policy_entropy) else policy_entropy, + 'wm_target_policy_entropy': average_target_policy_entropy.item(), + + + # ============ Step-wise Losses ============ + 'analysis/first_step_loss_value': first_step_losses.get('loss_value', torch.tensor(0.0)).item() if isinstance(first_step_losses.get('loss_value'), torch.Tensor) else 0.0, + 'analysis/first_step_loss_policy': first_step_losses.get('loss_policy', torch.tensor(0.0)).item() if isinstance(first_step_losses.get('loss_policy'), torch.Tensor) else 0.0, + 'analysis/first_step_loss_rewards': first_step_losses.get('loss_rewards', torch.tensor(0.0)).item() if isinstance(first_step_losses.get('loss_rewards'), torch.Tensor) else 0.0, + 'analysis/first_step_loss_obs': first_step_losses.get('loss_obs', torch.tensor(0.0)).item() if isinstance(first_step_losses.get('loss_obs'), torch.Tensor) else 0.0, + + 'analysis/middle_step_loss_value': middle_step_losses.get('loss_value', torch.tensor(0.0)).item() if isinstance(middle_step_losses.get('loss_value'), torch.Tensor) else 0.0, + 'analysis/middle_step_loss_policy': middle_step_losses.get('loss_policy', torch.tensor(0.0)).item() if isinstance(middle_step_losses.get('loss_policy'), torch.Tensor) else 0.0, + 'analysis/middle_step_loss_rewards': middle_step_losses.get('loss_rewards', torch.tensor(0.0)).item() if isinstance(middle_step_losses.get('loss_rewards'), torch.Tensor) else 0.0, + 'analysis/middle_step_loss_obs': middle_step_losses.get('loss_obs', torch.tensor(0.0)).item() if isinstance(middle_step_losses.get('loss_obs'), torch.Tensor) else 0.0, + + 'analysis/last_step_loss_value': last_step_losses.get('loss_value', torch.tensor(0.0)).item() if isinstance(last_step_losses.get('loss_value'), torch.Tensor) else 0.0, + 'analysis/last_step_loss_policy': last_step_losses.get('loss_policy', torch.tensor(0.0)).item() if isinstance(last_step_losses.get('loss_policy'), torch.Tensor) else 0.0, + 'analysis/last_step_loss_rewards': last_step_losses.get('loss_rewards', torch.tensor(0.0)).item() if isinstance(last_step_losses.get('loss_rewards'), torch.Tensor) else 0.0, + 'analysis/last_step_loss_obs': last_step_losses.get('loss_obs', torch.tensor(0.0)).item() if isinstance(last_step_losses.get('loss_obs'), torch.Tensor) else 0.0, + + # ============ Analysis Metrics ============ + 'analysis/latent_state_l2_norms': latent_state_l2_norms.item() if torch.is_tensor(latent_state_l2_norms) else latent_state_l2_norms, + 'analysis/latent_action_l2_norms': latent_action_l2_norms, + + # ============ Logits Statistics ============ + 'logits_value_mean': logits_value_mean, + 'logits_value_max': logits_value_max, + 'logits_value_min': logits_value_min, + 'logits_policy_mean': logits_policy_mean, + 'logits_policy_max': logits_policy_max, + 'logits_policy_min': logits_policy_min, + + # ============ Temperature Parameters ============ + 'temperature_value': temperature_value, + 'temperature_reward': temperature_reward, + 'temperature_policy': temperature_policy, + + # ============ Targets ============ + 'wm_target_reward': target_reward.mean().item(), + 'wm_target_value': target_value.mean().item(), + 'transformed_target_reward': transformed_target_reward.mean().item(), + 'transformed_target_value': transformed_target_value.mean().item(), + 'value_priority': value_priority_np.mean().item(), + 'value_priority_orig': value_priority_np, + + # ============ Gradient Norms ============ + 'wm_grad_norm': wm_grad_norm.item(), + + # ============ Learning Rates ============ + 'cur_lr_world_model': self._optimizer_world_model.param_groups[0]['lr'], + } + + return log_dict + + def _monitor_vars_learn(self) -> List[str]: + """ + Register variables to be monitored in learn mode. + These are logged by DI-engine's BaseLearner under "learner_iter/" prefix. + + Organized into groups: + - Core WM losses (essential for training diagnosis) + - WM analysis metrics (for deeper debugging) + - Training dynamics (LR, grad norm, entropy) + """ + return [ + # ---- Core WM Losses ---- + 'wm_total_loss', + 'wm_obs_loss', + 'wm_reward_loss', + 'wm_policy_loss', + 'wm_value_loss', + 'wm_latent_recon_loss', + 'wm_perceptual_loss', + + # ---- WM Policy Analysis ---- + 'wm_orig_policy_loss', + 'wm_policy_entropy', + 'wm_target_policy_entropy', + + # ---- WM Targets ---- + 'wm_target_reward', + 'wm_target_value', + 'value_priority', + + # ---- Adaptive Entropy ---- + 'adaptive_alpha', + 'adaptive_target_entropy_ratio', + 'alpha_loss', + + # ---- Training Dynamics ---- + 'wm_grad_norm', + 'cur_lr_world_model', + + # ---- Logits Statistics ---- + 'logits_value_mean', + 'logits_policy_mean', + + # ---- Temperature ---- + 'temperature_value', + 'temperature_reward', + 'temperature_policy', + + # ---- System ---- + 'Current_GPU', + 'Max_GPU', + ] + # ======================================================================== + + def pad_to_fixed_length(self, data, target_len=55, pad_val=-1e9, dtype=torch.float32): + """ + data: List[Sequence[Number]],每个元素长度可以不一样(比如 3 或 4) + 返回: tensor, 形状 [B, target_len],多余部分全是 pad_val + """ + batch_size = len(data) + out = torch.full((batch_size, target_len), pad_val, dtype=dtype) + for i, seq in enumerate(data): + if isinstance(seq, np.ndarray): + seq = seq.tolist() + L = min(len(seq), target_len) + if L > 0: + out[i, :L] = torch.tensor(seq[:L], dtype=dtype) + return out + + def _forward_collect( + self, + data: torch.Tensor, + action_mask: List[np.ndarray], + temperature: float = 1.0, + to_play: List[int] = None, + epsilon: float = 0.0, + ready_env_id: List[int] = None, + timestep: List = [0], + **kwargs + ) -> Dict[int, Dict[str, Any]]: + self._collect_model.eval() + + llm_prior_logprob = kwargs.pop('llm_prior_logprob', None) + valid_actions_list = kwargs.get('valid_actions_list', None) + current_envstep = kwargs.get('current_env_step', 0) + phase = kwargs.get('phase', None) + llm_collect_mode = kwargs.get('llm_collect_mode', None) + mcts_root_logits_dict = self.llm_cfg.mcts_root_logits_dict + + if llm_prior_logprob is None or not any(llm_prior_logprob) or mcts_root_logits_dict.mode == "wm_logits" or (phase == 'llm' and llm_collect_mode == 'wm_collect'): + logging.debug("No LLM priors provided, using standard UniZero MCTS") + return super()._forward_collect( + data, action_mask, temperature, to_play, epsilon, + ready_env_id=ready_env_id, timestep=timestep + ) + self._collect_mcts_temperature = temperature + self._collect_epsilon = epsilon + active_collect_env_num = data.shape[0] + if ready_env_id is None: + ready_env_id = np.arange(active_collect_env_num) + output = {i: None for i in ready_env_id} + + # Convert LLM priors to policy priors + # For Atari: llm_prior_logprob is a list of dicts (action_name -> prob) + # For text games: llm_prior_logprob is a list of dicts (action_name -> prob) + # Both use semantic action names now! + policy_priors = [] + for env_id in range(active_collect_env_num): + prior_data = llm_prior_logprob[env_id] + + # Check if this is a numpy array (legacy format) or dict (new format) + if isinstance(prior_data, np.ndarray): + # Legacy: numpy array with probabilities for each action index + prior = prior_data + elif isinstance(prior_data, dict): + # New format: dict mapping action names to probabilities + # Need to convert to array aligned with action space + actions = valid_actions_list[env_id] + prior = [] + + if len(actions) == 0: + # Fallback for edge case + print("Warning: No valid actions provided") + prior = np.ones(self.cfg.model.action_space_size) / self.cfg.model.action_space_size + else: + # Extract probabilities for each action in order + for action in actions: + prior.append(prior_data.get(action, 0.0)) + prior = np.array(prior, dtype=np.float32) + + # Normalize if needed + if prior.sum() > 0: + prior = prior / prior.sum() + else: + prior = np.ones(len(actions), dtype=np.float32) / len(actions) + else: + raise TypeError(f"Unexpected prior type: {type(prior_data)}") + + policy_priors.append(prior) + + policy_priors = self.pad_to_fixed_length(data=policy_priors, target_len=self.cfg.model.action_space_size, pad_val=-1e9) + + with torch.no_grad(): + network_output = self._collect_model.initial_inference(self.last_batch_obs_collect, self.last_batch_action_collect, data, timestep) + latent_state_roots, reward_roots, pred_values, policy_logits = mz_network_output_unpack(network_output) + + if mcts_root_logits_dict.mode == "llm_logits": + root_logits = policy_priors + + elif mcts_root_logits_dict.mode == "llm_plus_wm_logits": + llm_weight_min = mcts_root_logits_dict.llm_min_weight + llm_weight_max = mcts_root_logits_dict.llm_max_weight + + llm_probs = F.softmax(policy_priors, dim=-1) + mask_tensor = torch.from_numpy(np.stack(action_mask)) + policy_logits = policy_logits.cpu().masked_fill(mask_tensor == 0, -1e9) + wm_probs = F.softmax(policy_logits, dim=-1) + if mcts_root_logits_dict.plus_method == "adaptive": + wm_entropy = -(wm_probs * (wm_probs + 1e-8).log()).sum(dim=-1) + wm_entropy_norm = wm_entropy / torch.log(mask_tensor.sum(dim=-1).clamp(min=2.0)) + llm_weight = llm_weight_min + (llm_weight_max - llm_weight_min)*(1 - wm_entropy_norm) + combined_probs = (1 - llm_weight) * wm_probs + llm_probs * llm_weight + + elif mcts_root_logits_dict.plus_method == "fixed": + combined_probs = wm_probs * mcts_root_logits_dict.wm_weight + llm_probs * (1 - mcts_root_logits_dict.wm_weight) + + root_logits = torch.log(combined_probs + 1e-8) + + network_output.policy_logits = root_logits + if not self._cfg.mcts_ctree: + raise NotImplementedError("Python MCTS not supported for PriorZero") + + # ====================================================================== + # MCTS Search with LLM-Guided Priors + # ====================================================================== + pred_values_np = self.value_inverse_scalar_transform_handle(pred_values).detach().cpu().numpy() + latent_state_roots_np = latent_state_roots.detach().cpu().numpy() + policy_logits = root_logits.detach().cpu().numpy().tolist() + + legal_actions = [[i for i, x in enumerate(action_mask[j]) if x == 1] for j in range(active_collect_env_num)] + noises = [ + np.random.dirichlet([self._cfg.root_dirichlet_alpha] * int(sum(action_mask[j])) + ).astype(np.float32).tolist() for j in range(active_collect_env_num) + ] + roots = MCTSCtree.roots(active_collect_env_num, legal_actions) + roots.prepare(self._cfg.root_noise_weight, noises, reward_roots, policy_logits, to_play) + self._mcts_collect.search(roots, self._collect_model, latent_state_roots_np, to_play, timestep=timestep) + + roots_visit_count = roots.get_distributions() + roots_values = roots.get_values() + + batch_action = [] + for i, env_id in enumerate(ready_env_id): + distributions = roots_visit_count[i] + value = roots_values[i] + + action_index_in_legal_action_set, visit_count_distribution_entropy = select_action( + distributions, + temperature=self._collect_mcts_temperature, + deterministic=False + ) + + legal_action_indices = np.where(action_mask[i] == 1.0)[0] + action = legal_action_indices[action_index_in_legal_action_set] + + output[env_id] = { + 'action': int(action), + 'visit_count_distributions': distributions, + 'visit_count_distribution_entropy': visit_count_distribution_entropy, + 'searched_value': value, + 'predicted_value': pred_values_np[i], + 'predicted_policy_logits': policy_logits[i], + 'timestep': timestep[i], + } + if mcts_root_logits_dict.mode == "llm_plus_wm_logits": + if mcts_root_logits_dict.plus_method == "adaptive": + output[env_id]['llm_weight'] = llm_weight[i].item() + else: + output[env_id]['llm_weight'] = 1 - mcts_root_logits_dict.wm_weight + elif mcts_root_logits_dict.mode == "llm_logits": + output[env_id]['llm_weight'] = 1 + + batch_action.append(action) + self.last_batch_obs_collect = data + self.last_batch_action_collect = batch_action + return output + + def _forward_eval(self, data: torch.Tensor, action_mask: list, to_play: int = -1, + ready_env_id: np.array = None, timestep: List = [0], **kwargs) -> Dict: + self._eval_model.eval() + llm_prior_logprob = kwargs.pop('llm_prior_logprob', None) + valid_actions_list = kwargs.get('valid_actions_list', None) + mcts_root_logits_dict = self.llm_cfg.mcts_root_logits_dict + + if llm_prior_logprob is None or all(x is None for x in llm_prior_logprob) or mcts_root_logits_dict.mode == "wm_logits": + logging.debug("No LLM priors provided, using standard UniZero MCTS") + return super()._forward_eval( + data, action_mask, to_play=to_play, ready_env_id=ready_env_id, timestep=timestep + ) + + active_eval_env_num = data.shape[0] + if ready_env_id is None: + ready_env_id = np.arange(active_eval_env_num) + output = {i: None for i in ready_env_id} + mcts_info = {i: defaultdict(dict) for i in ready_env_id} + + policy_priors = [] + for env_id in range(active_eval_env_num): + prior_data = llm_prior_logprob[env_id] + actions = valid_actions_list[env_id] + + if isinstance(prior_data, np.ndarray): + # Image mode: numpy array with log-probs for each action index + prior = prior_data.tolist() + elif isinstance(prior_data, dict): + # Text mode: dict mapping action names to log-probs + prior = [] + if len(actions) == 0: + print("When valid actions is None, the action must be 'go'") + prior.append(prior_data['go']) + else: + for action in actions: + prior.append(prior_data[action]) + else: + # Fallback: uniform + prior = [0.0] * len(actions) if len(actions) > 0 else [0.0] + policy_priors.append(prior) + policy_priors = self.pad_to_fixed_length(data=policy_priors, target_len=self.cfg.model.action_space_size, pad_val=-1e9) + + with torch.no_grad(): + network_output = self._eval_model.initial_inference(self.last_batch_obs_eval, self.last_batch_action_eval, data, timestep) + latent_state_roots, reward_roots, pred_values, policy_logits = mz_network_output_unpack(network_output) + + if mcts_root_logits_dict.mode == "llm_logits": + root_logits = policy_priors + for env_id, prior, valid_actions in zip(ready_env_id, policy_priors, valid_actions_list): + llm_probs = F.softmax(prior, dim=-1).cpu().tolist() + for i in range(len(valid_actions)): + mcts_info[env_id]["root_llm_prob"][valid_actions[i]] = llm_probs[i] + + elif mcts_root_logits_dict.mode == "llm_plus_wm_logits": + llm_probs = F.softmax(policy_priors, dim=-1) + mask_tensor = torch.from_numpy(np.stack(action_mask)) + policy_logits = policy_logits.cpu().masked_fill(mask_tensor == 0, -1e9) + wm_probs = F.softmax(policy_logits, dim=-1) + combined_probs = wm_probs * mcts_root_logits_dict.wm_weight + llm_probs * (1 - mcts_root_logits_dict.wm_weight) + root_logits = torch.log(combined_probs + 1e-8) + + for env_id, llm_prob, wm_prob, combined_prob, valid_actions in zip(ready_env_id, llm_probs, wm_probs, combined_probs, valid_actions_list): + for i in range(len(valid_actions)): + if i < len(llm_prob) and i < len(wm_prob) and i < len(combined_prob): + mcts_info[env_id]["root_llm_prob"][valid_actions[i]] = llm_prob[i].item() + mcts_info[env_id]["root_wm_prob"][valid_actions[i]] = wm_prob[i].item() + mcts_info[env_id]["root_combined_prob"][valid_actions[i]] = combined_prob[i].item() + else: + break + + network_output.policy_logits = root_logits + + # if not in training, obtain the scalars of the value/reward + pred_values = self.value_inverse_scalar_transform_handle(pred_values).detach().cpu().numpy() # shape(B, 1) + latent_state_roots = latent_state_roots.detach().cpu().numpy() + policy_logits = root_logits.detach().cpu().numpy().tolist() + + legal_actions = [[i for i, x in enumerate(action_mask[j]) if x == 1] for j in range(active_eval_env_num)] + if self._cfg.mcts_ctree: + # cpp mcts_tree + roots = MCTSCtree.roots(active_eval_env_num, legal_actions) + else: + # python mcts_tree + roots = MCTSPtree.roots(active_eval_env_num, legal_actions) + roots.prepare_no_noise(reward_roots, policy_logits, to_play) + next_latent_state_with_env = self._mcts_eval.search(roots, self._eval_model, latent_state_roots, to_play, timestep) + + # list of list, shape: ``{list: batch_size} -> {list: action_space_size}`` + roots_visit_count_distributions = roots.get_distributions() + roots_values = roots.get_values() # shape: {list: batch_size} + + batch_action = [] + + for i, env_id in enumerate(ready_env_id): + distributions, value = roots_visit_count_distributions[i], roots_values[i] + # print("roots_visit_count_distributions:", distributions, "root_value:", value) + + # NOTE: Only legal actions possess visit counts, so the ``action_index_in_legal_action_set`` represents + # the index within the legal action set, rather than the index in the entire action set. + # Setting deterministic=True implies choosing the action with the highest value (argmax) rather than + # sampling during the evaluation phase. + action_index_in_legal_action_set, visit_count_distribution_entropy = select_action( + distributions, temperature=1, deterministic=True + ) + # NOTE: Convert the ``action_index_in_legal_action_set`` to the corresponding ``action`` in the + # entire action set. + action = np.where(action_mask[i] == 1.0)[0][action_index_in_legal_action_set] + + # Predict the next latent state based on the selected action and policy + next_latent_state = next_latent_state_with_env[i][action] + + output[env_id] = { + 'action': action, + 'visit_count_distributions': distributions, + 'visit_count_distribution_entropy': visit_count_distribution_entropy, + 'searched_value': value, + 'predicted_value': pred_values[i], + 'predicted_policy_logits': policy_logits[i], + 'timestep': timestep[i], + } + batch_action.append(action) + for idx, action in enumerate(valid_actions_list[i]): + if idx < len(distributions): + mcts_info[env_id]["visit_count_distributions"][action] = distributions[idx] + else: + break + self.last_batch_obs_eval = data + self.last_batch_action_eval = batch_action + + return output, mcts_info diff --git a/zoo/jericho/priorzero/src/priorzero_trainer.py b/zoo/jericho/priorzero/src/priorzero_trainer.py new file mode 100644 index 000000000..b3b4e8d56 --- /dev/null +++ b/zoo/jericho/priorzero/src/priorzero_trainer.py @@ -0,0 +1,222 @@ +from __future__ import annotations +import hashlib +import os +import copy +import json +import logging + +from typing import Any, Dict, List, Optional, Tuple + +import torch +import torch.nn.functional as F +import ray +import numpy as np +from transformers import AutoTokenizer +import torch.distributed as dist + +import ray +import torch + +import numpy as np + + +class AdaptiveKLController: + """ + Adaptive KL controller described in the paper: + https://arxiv.org/pdf/1909.08593.pdf + """ + + def __init__(self, init_kl_coef, target, horizon): + self.value = init_kl_coef + self.target = target + self.horizon = horizon + + def update(self, current, n_steps): + target = self.target + proportional_error = np.clip(current / target - 1, -0.2, 0.2) + mult = 1 + proportional_error * n_steps / self.horizon + self.value *= mult + + +class FixedKLController: + """Fixed KL controller.""" + + def __init__(self, kl_coef): + self.value = kl_coef + + def update(self, current, n_steps): + pass + + +def get_tokenizer(pretrain: str) -> AutoTokenizer: + tokenizer = AutoTokenizer.from_pretrained( + pretrain, trust_remote_code=True, padding_side="left" + ) + if tokenizer.pad_token is None: + tokenizer.pad_token = tokenizer.eos_token + return tokenizer + +class PriorZeroLLMTrainer: + + def __init__( + self, + cfg, + pretrain: str, + strategy, + vllm_engine, + policy_model, # RayActorGroup(PolicyModelActor) + reference_model=None, # RayActorGroup(ReferenceModelActor) or None + exp_name: str = None, + tb_logger = None, + instance_name: str = "llm_ppo", + llm_save_freq: int = 1000, + ): + self.cfg = cfg + self.pretrain = pretrain + self.strategy = strategy + self.args = getattr(strategy, "args", None) + + self.policy_model = policy_model + self.reference_model = reference_model + self.vllm_engine = vllm_engine + self.global_step = 0 + self.llm_save_freq = llm_save_freq + + self.tokenizer = get_tokenizer(self.pretrain) + + self.init_kl_coef = float(getattr(cfg, "rft_kl_coef", 0.0)) + + self.kl_ctl = FixedKLController(self.init_kl_coef) + self.rank = self.strategy.get_rank() + self.world_size = self.strategy.world_size + self.instance_name = instance_name + self._tb_prefix = instance_name.replace("_ppo", "") # e.g. "llm" or "vl" + + if tb_logger is not None: + from ding.utils import build_logger + self._logger, _ = build_logger( + path=f'./{exp_name}/log/{instance_name}', name=instance_name, need_tb=False + ) + self._tb_logger = tb_logger + else: + self._logger = None + self._tb_logger = None + + def train_batch(self, data, collect_env_steps) -> Dict[str, float]: + if data is None: + return {} + input_ids, attention_mask, action_mask, advantage, rollout_lp, log_status = data + assert len(input_ids) == len(attention_mask) == len(action_mask) == len(advantage) == len(rollout_lp) == len(log_status) + batch_input_stats = self._collect_input_ids_stats(input_ids) + + batch = { + "input_ids": input_ids, + "attention_mask": attention_mask, + "action_mask": action_mask, + "advantages": advantage, + "rollout_action_logprob": rollout_lp, + "log_status": log_status, + } + if self.reference_model is not None: + base_action_log_probs = self.reference_model.forward( + sequences = batch['input_ids'], + action_mask = batch['action_mask'], + attention_mask=batch['attention_mask'], + ) + batch["ref_action_log_probs"] = base_action_log_probs + else: + batch["ref_action_log_probs"] = None + + if self.strategy.args.deepspeed_enable_sleep: + self.policy_model.reload_states() + + old_action_log_probs = self.policy_model.forward( + sequences = batch['input_ids'], + action_mask = batch['action_mask'], + attention_mask=batch['attention_mask'], + ) + batch["old_action_log_probs"] = old_action_log_probs + + status = self.policy_model.fit(batch, self.kl_ctl) + + if self.strategy.args.deepspeed_enable_sleep: + self.policy_model.offload_states() + + if self.vllm_engine is not None: + self._broadcast_to_vllm() + + for tmp_dict in status: + tmp_dict.update(batch_input_stats) + + if self._tb_logger is not None and self.strategy.is_rank_0(): + logging.getLogger("priorzero.train").info( + f"[LLM] samples={int(batch_input_stats['input_ids_global_sample_count'])} " + f"unique={int(batch_input_stats['input_ids_global_unique_count'])} " + f"ratio={float(batch_input_stats['input_ids_global_unique_ratio']):.4f}" + ) + for tmp_dict in status: + for k, v in tmp_dict.items(): + if k == 'iter': + continue + self._tb_logger.add_scalar(f"learner_{self._tb_prefix}_iter/{k}", float(v), int(tmp_dict['iter'])) + self._tb_logger.add_scalar(f"learner_{self._tb_prefix}_envstep/{k}", float(v), int(collect_env_steps)) + self.global_step = max(self.global_step, int(tmp_dict['iter'])) + + self._sync_global_step_from_rank0() + + if self.global_step > 0 and self.global_step % self.llm_save_freq == 0: + self.policy_model.save_model() + + def get_state(self) -> Dict[str, Any]: + kl_val = float(self.kl_ctl.value) if hasattr(self.kl_ctl, "value") else float(self.init_kl_coef) + return {"global_step": self.global_step, "kl_coef": kl_val} + + def _sync_global_step_from_rank0(self): + if self.world_size <= 1: + return + lst = [self.global_step] if self.rank == 0 else [None] + dist.broadcast_object_list(lst, src=0) + self.global_step = int(lst[0]) + + def _collect_input_ids_stats(self, input_ids: torch.Tensor) -> Dict[str, float]: + local_hashes = self._hash_input_rows(input_ids) + local_sample_count = len(local_hashes) + local_unique_count = len(set(local_hashes)) + + global_hashes = local_hashes + if self.world_size > 1: + gathered_hashes = [None for _ in range(self.world_size)] + dist.all_gather_object(gathered_hashes, local_hashes) + global_hashes = [item for rank_hashes in gathered_hashes for item in rank_hashes] + + global_sample_count = len(global_hashes) + global_unique_count = len(set(global_hashes)) + global_duplicate_count = global_sample_count - global_unique_count + global_unique_ratio = global_unique_count / global_sample_count if global_sample_count > 0 else 0.0 + + return { + "input_ids_local_sample_count": float(local_sample_count), + "input_ids_local_unique_count": float(local_unique_count), + "input_ids_global_sample_count": float(global_sample_count), + "input_ids_global_unique_count": float(global_unique_count), + "input_ids_global_duplicate_count": float(global_duplicate_count), + "input_ids_global_unique_ratio": float(global_unique_ratio), + } + + def _hash_input_rows(self, input_ids: torch.Tensor) -> List[str]: + input_ids_cpu = input_ids.detach().to("cpu") + return [ + hashlib.blake2b(row.numpy().tobytes(), digest_size=16).hexdigest() + for row in input_ids_cpu + ] + + def _broadcast_to_vllm(self): + if self.strategy.args.vllm_enable_sleep: + self.vllm_engine.wake_up() + + logging.getLogger("priorzero.train").info("[LLM] vLLM weight sync start") + self.policy_model.broadcast_to_vllm() + logging.getLogger("priorzero.train").info("[LLM] vLLM weight sync done") + + if self.strategy.args.vllm_enable_sleep: + self.vllm_engine.sleep() \ No newline at end of file diff --git a/zoo/jericho/priorzero/src/strategy/deepspeed.py b/zoo/jericho/priorzero/src/strategy/deepspeed.py new file mode 100644 index 000000000..0b69c8529 --- /dev/null +++ b/zoo/jericho/priorzero/src/strategy/deepspeed.py @@ -0,0 +1,644 @@ +import os +import shutil +from abc import ABC +from collections import defaultdict +from datetime import timedelta +from typing import List, Tuple, Union +import math + +import deepspeed +import torch +import torch.nn as nn +import torch.optim as optim +import transformers +from deepspeed.ops.adam import DeepSpeedCPUAdam, FusedAdam +from peft import PeftModel, get_peft_model_state_dict +from torch import distributed as dist +from torch.distributed.device_mesh import init_device_mesh +from torch.optim import Optimizer + +from utils import torch_dist_barrier_and_cuda_sync +from models.actor import Actor +from packaging import version + +ModelOptimPair = Tuple[nn.Module, Optimizer] +ModelOrModelOptimPair = Union[nn.Module, ModelOptimPair] + + +def get_train_ds_config( + offload, + adam_offload=True, + stage=2, + bf16=True, + max_norm=1.0, + zpg=8, + grad_accum_dtype=None, + overlap_comm=False, + use_ds_universal_ckpt=False, + deepcompile=False, + tensor_parallel_size=1, +): + device = "cpu" if offload else "none" + zero_opt_dict = { + "stage": stage, + "offload_param": {"device": device}, + "offload_optimizer": { + "device": "cpu" if adam_offload else "none", + "pin_memory": True, + }, + "sub_group_size": "auto", + "stage3_max_live_parameters": "auto", + "stage3_max_reuse_distance": "auto", + "stage3_param_persistence_threshold": "auto", + "stage3_prefetch_bucket_size": "auto", + "reduce_bucket_size": "auto", + # ZeRO++ + "zero_hpz_partition_size": zpg, + "zero_quantized_weights": False, + "zero_quantized_gradients": False, + } + if overlap_comm: + zero_opt_dict["overlap_comm"] = True + zero_opt_dict["contiguous_gradients"] = True + if stage == 3: + zero_opt_dict["reduce_scatter"] = True + + return { + "steps_per_print": 100, + "zero_optimization": zero_opt_dict, + "bf16": { + "enabled": bf16, + }, + "gradient_clipping": max_norm, + "prescale_gradients": False, + "wall_clock_breakdown": False, + "data_types": {"grad_accum_dtype": grad_accum_dtype}, + "checkpoint": { + "load_universal": use_ds_universal_ckpt, + }, + "compile": { + "deepcompile": deepcompile, + }, + "tensor_parallel": { + "autotp_size": tensor_parallel_size, + }, + } + + +def get_eval_ds_config( + offload, + stage=0, + bf16=True, + deepcompile=False, + tensor_parallel_size=1, +): + # At least for 0.16.6, DeepCompile hasn't support pure inference mode + # https://github.com/deepspeedai/DeepSpeed/pull/7225 + deepcompile = False + + zero_opt_dict = { + "stage": stage, + "stage3_max_live_parameters": "auto", + "stage3_max_reuse_distance": "auto", + "stage3_param_persistence_threshold": "auto", + "stage3_prefetch_bucket_size": "auto", + "offload_param": { + "device": "cpu" if offload else "none", + "pin_memory": True, + }, + } + return { + "steps_per_print": 100, + "zero_optimization": zero_opt_dict, + "bf16": { + "enabled": bf16, + }, + "gradient_clipping": 1.0, + "prescale_gradients": False, + "wall_clock_breakdown": False, + "compile": { + "deepcompile": deepcompile, + }, + "tensor_parallel": { + "autotp_size": tensor_parallel_size, + }, + } + + +def get_optimizer_grouped_parameters( + model, + weight_decay, + no_decay_name_list=["bias", "layer_norm.weight", "layernorm.weight", "norm.weight", "ln_f.weight"], +): + optimizer_grouped_parameters = [ + { + "params": [ + p + for n, p in model.named_parameters() + if (not any(nd in n for nd in no_decay_name_list) and p.requires_grad) + ], + "weight_decay": weight_decay, + }, + { + "params": [ + p + for n, p in model.named_parameters() + if (any(nd in n for nd in no_decay_name_list) and p.requires_grad) + ], + "weight_decay": 0.0, + }, + ] + return optimizer_grouped_parameters + +def offload_deepspeed_states(model, pin_memory=True, non_blocking=True): + zero_stage = model.zero_optimization_stage() # config['zero_optimization']['stage'] + adam_offload = model.config["zero_optimization"]["offload_optimizer"]["device"] == "cpu" + + # state offloading not required when using Adam optimizer offloading + if adam_offload: + return + + if zero_stage != 3 and version.parse(deepspeed.__version__) <= version.parse("0.17.5"): + raise NotImplementedError( + "Only Zero stage 3 is currently supported when using DeepSpeed version 0.17.5 or lower" + ) + + # if zero_stage == 3 and not adam_offload: + from deepspeed.runtime.zero.offload_config import OffloadDeviceEnum, OffloadStateTypeEnum + + offload_state_types = [ + OffloadStateTypeEnum.optim_states, + OffloadStateTypeEnum.contiguous_grad_buffer, + OffloadStateTypeEnum.hp_params, + ] + + if version.parse(deepspeed.__version__) >= version.parse("0.16.5"): + # These offload types are fixed in https://github.com/deepspeedai/DeepSpeed/pull/7050 + offload_state_types += [ + OffloadStateTypeEnum.lp_grads, + # OffloadStateTypeEnum.lp_params, + ] + + model.optimizer.offload_states( + include=offload_state_types, + device=OffloadDeviceEnum.cpu, + pin_memory=pin_memory, + non_blocking=non_blocking, + ) + model.empty_partition_cache() + torch.cuda.empty_cache() + torch.distributed.barrier() + torch.cuda.synchronize() + +def reload_deepspeed_states(model, non_blocking=True): + zero_stage = model.zero_optimization_stage() # config['zero_optimization']['stage'] + adam_offload = model.config["zero_optimization"]["offload_optimizer"]["device"] == "cpu" + + # state offloading not required when using Adam optimizer offloading + if adam_offload: + return + + if zero_stage != 3 and version.parse(deepspeed.__version__) <= version.parse("0.17.5"): + raise NotImplementedError( + "Only Zero stage 3 is currently supported when using DeepSpeed version 0.17.5 or lower" + ) + model.reload_states(non_blocking=non_blocking) + torch.cuda.empty_cache() + torch.distributed.barrier() + torch.cuda.synchronize() + +from deepspeed.runtime.zero.partition_parameters import ZeroParamStatus +def _z3_params_to_fetch(param_list): + return [p for p in param_list if hasattr(p, "ds_id") and p.ds_status == ZeroParamStatus.NOT_AVAILABLE] + + +def get_strategy(args): + strategy = DeepspeedStrategy( + seed=getattr(args, "seed", 42), + max_norm=getattr(args, "max_norm", 1.0), + micro_train_batch_size=getattr(args, "micro_train_batch_size", 1), + train_batch_size=getattr(args, "train_batch_size", 128), + zero_stage=args.zero_stage, + bf16=getattr(args, "bf16", True), + args=args, + ) + return strategy + + +class DeepspeedStrategy(ABC): + """ + The strategy for training with Accelerator. + """ + + def __init__( + self, + seed: int = 42, + max_norm: float = 0.0, + micro_train_batch_size=1, + train_batch_size=1, + zero_stage=2, + bf16=True, + args=None, + ) -> None: + super().__init__() + + self.args = args + self.stage = zero_stage + self.train_batch_size = train_batch_size + self.micro_train_batch_size = micro_train_batch_size + self.bf16 = bf16 + self.seed = seed + self.max_norm = max_norm + + self.adam_offload = getattr(args, "adam_offload", False) + self.zpg = getattr(args, "zpg", 1) + self.grad_accum_dtype = getattr(args, "grad_accum_dtype", None) + self.overlap_comm = getattr(args, "overlap_comm", False) + self.deepcompile = getattr(args, "deepcompile", False) + self.ds_tensor_parallel_size = getattr(args, "ds_tensor_parallel_size", 1) + self.use_dynamic_batch = getattr(self.args, "use_dynamic_batch", False) + + if self.ds_tensor_parallel_size > 1: + assert deepspeed.version >= "0.16.4", "DeepSpeed version must be >= 0.16.4 for tensor parallel training" + assert bf16, "BF16 is required for tensor parallel training" + + self.is_rlhf = False + self.time_steps = defaultdict(int) + + def setup_distributed(self, timeout=timedelta(minutes=60)) -> None: + transformers.set_seed(self.seed) + + local_rank = int(os.environ.get("LOCAL_RANK", "-1")) + if local_rank != -1: + torch.cuda.set_device(local_rank) + + # Initializes the distributed backend which will take care of synchronizing nodes/GPUs + # deepspeed.init_distributed(dist_backend="nccl", timeout=timeout) + + if not dist.is_initialized(): + dist.init_process_group(backend="nccl", timeout=timeout) + + # mesh + self.world_size = dist.get_world_size() + dp_size = self.world_size // self.ds_tensor_parallel_size + self.ds_device_mesh = init_device_mesh( + "cuda", (dp_size, self.ds_tensor_parallel_size), mesh_dim_names=("dp", "tp") + ) + + self.accumulated_gradient = ( + self.train_batch_size + * self.ds_tensor_parallel_size + // self.micro_train_batch_size + // self.world_size + ) + + def create_optimizer(self, model, **kwargs) -> Optimizer: + if isinstance(model, Actor): + model = model.model + # Optimizer + AdamOptimizer = DeepSpeedCPUAdam if self.adam_offload else FusedAdam + optim_params = get_optimizer_grouped_parameters(model, kwargs["weight_decay"]) + optim = AdamOptimizer(optim_params, **kwargs) + return optim + + def backward(self, loss: torch.Tensor, model: nn.Module, optimizer: optim.Optimizer, **kwargs) -> None: + if isinstance(model, Actor): + model = model.model + model.backward(loss) + + def optimizer_step( + self, + optimizer: optim.Optimizer, + model: nn.Module, + scheduler, + name="model", + **kwargs, + ) -> None: + if isinstance(model, Actor): + model = model.model + model.step() + + + def _unwrap_model(self, model) -> nn.Module: + if isinstance(model, Actor): + return self._unwrap_model(model.model) + elif hasattr(model, "module"): + return model.module + else: + return model + + def prepare( + self, *models_or_model_optim_pairs: ModelOrModelOptimPair, is_rlhf=False + ) -> Union[List[ModelOrModelOptimPair], ModelOrModelOptimPair]: + ret = [] + self.is_rlhf = is_rlhf + for arg in models_or_model_optim_pairs: + if isinstance(arg, tuple): + assert len(arg) == 3, f'Expect (model, optimizer, scheduler) pair, got a tuple with size "{len(arg)}"' + if arg[0] is not None: + ret.append(self._ds_init_train_model(*arg)) + else: + ret.append((None, None, None)) + else: + ret.append(self._ds_init_eval_model(arg)) + + return ret[0] if len(ret) == 1 else ret + + def _ds_init_train_model(self, model, optim, scheduler): + is_actor = isinstance(model, Actor) + ds_config = self.get_ds_train_config(is_actor) + + if self.ds_tensor_parallel_size > 1: + tp_model = deepspeed.tp_model_init( + model=model.model if is_actor else model, tp_size=self.ds_tensor_parallel_size, dtype=torch.bfloat16 + ) + if is_actor: + model.model = tp_model + else: + model = tp_model + + engine, optim, _, scheduler = deepspeed.initialize( + model=model.model if is_actor else model, + optimizer=optim, + lr_scheduler=scheduler, + config=ds_config, + args={"local_rank": int(os.environ.get("LOCAL_RANK", "-1"))}, + dist_init_required=True, + ) + if self.deepcompile: + engine.compile() + if is_actor: + model.model = engine + else: + model = engine + + return model, optim, scheduler + + def get_ds_train_config(self, is_actor): + # DS Config + ds_config = get_train_ds_config( + offload=False, + adam_offload=self.adam_offload, + stage=self.stage, + bf16=self.bf16, + max_norm=self.max_norm, + zpg=self.zpg, + grad_accum_dtype=self.grad_accum_dtype, + overlap_comm=self.overlap_comm, + deepcompile=self.deepcompile, + tensor_parallel_size=self.ds_tensor_parallel_size, + ) + if self.use_dynamic_batch: + ds_config["train_micro_batch_size_per_gpu"] = 1 + ds_config["gradient_accumulation_steps"] = 1 + else: + ds_config["train_micro_batch_size_per_gpu"] = self.micro_train_batch_size + ds_config["train_batch_size"] = self.train_batch_size * self.ds_tensor_parallel_size + + return ds_config + + def _ds_init_eval_model(self, model): + if not model: + return model + is_actor = isinstance(model, Actor) + ds_config = self.get_ds_eval_config(offload=getattr(model, "_offload", False)) + + if self.ds_tensor_parallel_size > 1: + tp_model = deepspeed.tp_model_init( + model=model.model if is_actor else model, tp_size=self.ds_tensor_parallel_size, dtype=torch.bfloat16 + ) + if is_actor: + model.model = tp_model + else: + model = tp_model + + engine, *_ = deepspeed.initialize( + model=model.model if is_actor else model, + args={"local_rank": int(os.environ.get("LOCAL_RANK", "-1"))}, + config=ds_config, + dist_init_required=True, + ) + if self.deepcompile: + engine.compile() + if is_actor: + model.model = engine + else: + model = engine + return model + + def get_ds_eval_config(self, offload=False): + # DS Config + ds_config = get_eval_ds_config( + offload=offload, + stage=self.stage if self.stage == 3 else 0, + bf16=self.bf16, + deepcompile=self.deepcompile, + tensor_parallel_size=self.ds_tensor_parallel_size, + ) + ds_config["train_micro_batch_size_per_gpu"] = self.micro_train_batch_size + ds_config["train_batch_size"] = self.train_batch_size * self.ds_tensor_parallel_size + + return ds_config + + def moving_average(self, model, model_ema, beta=0.992, device="cpu"): + self.time_steps["ema"] += 1 + if self.time_steps["ema"] % self.accumulated_gradient == 0 or self.use_dynamic_batch: + with torch.no_grad(): + for param, param_ema in zip(model.parameters(), model_ema.parameters()): + if param.requires_grad: + if self.stage != 3: + data = param.data.to(device) + param_ema.data.copy_((1 - beta) * data + beta * param_ema.data) + else: + # TODO: use prefiltering for efficiency + params_to_fetch = _z3_params_to_fetch([param, param_ema]) + with deepspeed.zero.GatheredParameters(params_to_fetch, enabled=len(params_to_fetch) > 0): + data = param.data.to(device) + param_ema.data.copy_((1 - beta) * data + beta * param_ema.data) + + def load_model( + self, + model: nn.Module, + path: str, + map_location="cpu", + strict: bool = False, + key_replace_fn=None, + ) -> None: + unwrapped_model = self._unwrap_model(model) + state_dict = torch.load(path, map_location=map_location) + if key_replace_fn: + state_dict = key_replace_fn(state_dict) + unwrapped_model.load_state_dict(state_dict, strict=strict) + + def save_model(self, model: nn.Module, tokenizer, output_dir, **kwargs) -> None: + if self.is_rank_0(): + os.makedirs(output_dir, exist_ok=True) + + # save model weights for ZeRO2/3 + model_to_save = self._unwrap_model(model) + + # gather parameters + if self.args.zero_stage > 2 or self.args.ds_tensor_parallel_size > 1: + output_state_dict = ( + model.model._consolidated_16bit_state_dict() + if isinstance(model, Actor) + else model._consolidated_16bit_state_dict() + ) + else: + from deepspeed.checkpoint.utils import clone_tensors_for_torch_save + + output_state_dict = clone_tensors_for_torch_save(model_to_save.state_dict()) + + if self.is_rank_0(): + state_dict_keys = set(model_to_save.state_dict().keys()) + output_state_dict_keys = set(output_state_dict.keys()) + + # corner case for tie_word_embeddings, such as Qwen2-0.5B + if getattr(model_to_save.config, "tie_word_embeddings", False) and "lm_head.weight" in state_dict_keys: + state_dict_keys.remove("lm_head.weight") + + assert state_dict_keys.issubset( + output_state_dict_keys + ), f"mismatch keys {output_state_dict_keys.symmetric_difference(state_dict_keys)}" + + # only save peft weights https://github.com/microsoft/DeepSpeed/issues/4295 + if isinstance(model_to_save, PeftModel): + model_to_save.save_pretrained(output_dir, **kwargs) + if self.ds_tensor_parallel_size > 1 or self.stage == 3: + torch.save( + get_peft_model_state_dict(model_to_save, output_state_dict), + os.path.join(output_dir, "adapter_model.bin"), + ) + filename = os.path.join(output_dir, "adapter_model.safetensors") + if os.path.exists(filename): + os.remove(filename) + else: + # save model + model_to_save.save_pretrained(output_dir, state_dict=output_state_dict, **kwargs) + + # save config + output_config_file = os.path.join(output_dir, "config.json") + model_to_save.config.to_json_file(output_config_file) + # save tokenizer + tokenizer.save_pretrained(output_dir) + + del output_state_dict + # Explicitly release memory + import gc + + gc.collect() + + torch_dist_barrier_and_cuda_sync() + + def all_reduce(self, data, op="mean"): + assert op in ("mean", "max", "sum") + if isinstance(data, dict): + ret = {} + for k, v in data.items(): + ret[k] = self.all_reduce(v, op) + return ret + else: + is_tensor = True + if not isinstance(data, torch.Tensor): + data = torch.Tensor([data]) + is_tensor = False + is_cpu_tensor = data.device.type == "cpu" + + if is_cpu_tensor: + data = data.to(torch.cuda.current_device()) + if op == "mean": + data /= self.world_size + dist.all_reduce(data, op=dist.ReduceOp.MAX if op == "max" else dist.ReduceOp.SUM) + if is_cpu_tensor: + data = data.cpu() + return data.item() if not is_tensor else data + + def all_gather(self, data): + if isinstance(data, dict): + ret = {} + for k, v in data.items(): + ret[k] = self.all_gather(v) + return ret + else: + if not isinstance(data, torch.Tensor): + data = torch.Tensor([data]) + is_cpu_tensor = data.device.type == "cpu" + + ret = [torch.zeros_like(data).to(torch.cuda.current_device()) for _ in range(self.world_size)] + dist.all_gather(ret, data.to(torch.cuda.current_device())) + return torch.cat(ret).cpu() if is_cpu_tensor else torch.cat(ret) + + def print(self, *msg): + if self.is_rank_0(): + print(*msg) + + def is_rank_0(self) -> bool: + if not dist.is_initialized(): + return True + return dist.get_rank() == 0 + + def get_rank(self) -> int: + if not dist.is_initialized(): + return 0 + return dist.get_rank() + + def save_ckpt(self, model, save_dir, tag=None, max_num=3, max_mem=1000, client_state={}, save_latest=True): + assert isinstance(model, deepspeed.DeepSpeedEngine) + if self.is_rank_0(): + os.makedirs(save_dir, exist_ok=True) + MAX_SIZE = max_mem * 1024**3 # Convert GB to bytes + + while True: + subdirs = sorted( + [ + (os.path.join(save_dir, d), os.path.getmtime(os.path.join(save_dir, d))) + for d in os.listdir(save_dir) + if os.path.isdir(os.path.join(save_dir, d)) + ], + key=lambda x: x[1], + ) + total_size = sum( + os.path.getsize(os.path.join(dirpath, f)) + for subdir, _ in subdirs + for dirpath, _, filenames in os.walk(subdir) + for f in filenames + ) + + if len(subdirs) >= max_num or total_size > MAX_SIZE: + oldest_dir = subdirs[0][0] + if os.path.exists(oldest_dir): + shutil.rmtree(oldest_dir) + self.print(f"Deleted oldest ckpt {oldest_dir}") + else: + break + + torch_dist_barrier_and_cuda_sync() + model.save_checkpoint(save_dir, tag=tag, client_state=client_state, save_latest=save_latest) + + # Explicitly release memory + import gc + + gc.collect() + + def load_ckpt( + self, + model, + load_dir, + tag=None, + load_module_strict=True, + load_optimizer_states=True, + load_lr_scheduler_states=True, + load_module_only=False, + ): + assert isinstance(model, deepspeed.DeepSpeedEngine) + load_path, states = model.load_checkpoint( + load_dir, + tag, + load_module_strict=load_module_strict, + load_optimizer_states=load_optimizer_states, + load_lr_scheduler_states=load_lr_scheduler_states, + load_module_only=load_module_only, + ) + if load_path is None: + raise Exception(f"[deepspeed] failed to resume from checkpoint {load_dir}") + return load_path, states diff --git a/zoo/jericho/priorzero/src/utils.py b/zoo/jericho/priorzero/src/utils.py new file mode 100644 index 000000000..6377097fe --- /dev/null +++ b/zoo/jericho/priorzero/src/utils.py @@ -0,0 +1,230 @@ +import torch +import torch.nn.functional as F +from typing import List, Dict, Any, Tuple, Union, Optional +from transformers import AutoTokenizer +from dataclasses import is_dataclass +import os +import logging +import inspect +import textwrap + + +# ============================================================================ +# Structured Logging Setup +# ============================================================================ + +def setup_priorzero_logging(exp_name: str, rank: int = 0) -> Dict[str, logging.Logger]: + """ + Create structured loggers for PriorZero training. + Only rank 0 gets console output and file handlers. + Other ranks get NullHandler (silent). + + Returns dict with keys: 'main', 'train', 'eval' + """ + log_dir = os.path.join(exp_name, "run_logs") + os.makedirs(log_dir, exist_ok=True) + + file_fmt = logging.Formatter("%(asctime)s [%(levelname)s] %(message)s", datefmt="%Y-%m-%d %H:%M:%S") + console_fmt = logging.Formatter("[%(levelname).1s] %(message)s") + + loggers = {} + for name, filename in [("main", "main.log"), ("train", "train.log"), ("eval", "eval.log")]: + lg = logging.getLogger(f"priorzero.{name}") + lg.setLevel(logging.DEBUG) + lg.handlers.clear() + lg.propagate = False + + if rank == 0: + fh = logging.FileHandler(os.path.join(log_dir, filename), mode="a") + fh.setLevel(logging.DEBUG) + fh.setFormatter(file_fmt) + lg.addHandler(fh) + + ch = logging.StreamHandler() + ch.setLevel(logging.INFO) + ch.setFormatter(console_fmt) + lg.addHandler(ch) + else: + lg.addHandler(logging.NullHandler()) + + loggers[name] = lg + + # Error log: captures WARNING+ from all priorzero loggers + if rank == 0: + err_handler = logging.FileHandler(os.path.join(log_dir, "error.log"), mode="a") + err_handler.setLevel(logging.WARNING) + err_handler.setFormatter(file_fmt) + for lg in loggers.values(): + lg.addHandler(err_handler) + + return loggers + +def dump_dataclass_cfg_py(cfg, path: str) -> str: + if not is_dataclass(cfg): + raise TypeError(type(cfg)) + + def norm(x): + if isinstance(x, dict): + return {k: norm(v) for k, v in x.items()} + if hasattr(x, "__class__") and x.__class__.__name__ == "EasyDict": + return {k: norm(v) for k, v in dict(x).items()} + if isinstance(x, (list, tuple)): + t = [norm(v) for v in x] + return tuple(t) if isinstance(x, tuple) else t + return x + cls = type(cfg) + fields = cls.__dataclass_fields__.keys() + lines = [f"{k} = {repr(norm(getattr(cfg, k)))}" for k in fields] + [""] + with open(path, "w", encoding="utf-8") as f: + f.write("\n".join(lines)) + return + +def torch_dist_barrier_and_cuda_sync(): + """Synchronize distributed training and CUDA operations. + This function ensures that: + 1. All distributed processes reach this point (barrier) + 2. All CUDA operations are completed (synchronize) + """ + import torch + + torch.distributed.barrier() + torch.cuda.synchronize() + + +def get_tokenizer(pretrain, model, padding_side="left", use_fast=True): + tokenizer = AutoTokenizer.from_pretrained(pretrain, trust_remote_code=True, use_fast=use_fast) + tokenizer.padding_side = padding_side + if tokenizer.pad_token is None: + tokenizer.pad_token = tokenizer.eos_token + tokenizer.pad_token_id = tokenizer.eos_token_id + if model is not None: + model.config.pad_token_id = tokenizer.pad_token_id + + return tokenizer + +@torch.compile +def compute_entropy(logits: torch.Tensor): + pd = torch.nn.functional.softmax(logits, dim=-1) + entropy = torch.logsumexp(logits, dim=-1) - torch.sum(pd * logits, dim=-1) + return entropy + + +def compute_approx_kl( + log_probs: torch.Tensor, + log_probs_base: torch.Tensor, + kl_estimator: str = "k1", +) -> torch.Tensor: + """ + Compute the approximate KL divergence between two distributions. + Schulman blog: http://joschu.net/blog/kl-approx.html + + Args: + log_probs: Log probabilities of the new distribution. + log_probs_base: Log probabilities of the base distribution. + """ + + if kl_estimator == "k1": + log_ratio = log_probs.float() - log_probs_base.float() + + # The k2 estimator is the non negative kl approximation in + # http://joschu.net/blog/kl-approx.html + # The k2_loss is approximately equivalent to the + # one-step KL divergence penalty with the k1 estimator + # used in https://arxiv.org/pdf/2310.10505. + if kl_estimator == "k2": + log_ratio = log_probs.float() - log_probs_base.float() + log_ratio = log_ratio**2 / 2.0 + + # The k3 estimator is the non negative kl approximation in + # http://joschu.net/blog/kl-approx.html + if kl_estimator == "k3": + log_ratio = log_probs.float() - log_probs_base.float() + log_ratio = -log_ratio + log_ratio = log_ratio.exp() - 1 - log_ratio + + log_ratio = log_ratio.clamp(min=-10, max=10) + return log_ratio + +def masked_mean(tensor: torch.Tensor, mask: Optional[torch.Tensor], dim: int = None) -> torch.Tensor: + if mask is None: + return tensor.mean(dim=dim) + return (tensor * mask).sum(dim=dim) / mask.sum(dim=dim) + + +def _logsumexp_by_chunk(logits: torch.Tensor, chunk_size: int = 1024) -> torch.Tensor: + seq_len = logits.shape[0] + logsumexp_values = torch.zeros((seq_len), device=logits.device, dtype=logits.dtype) + for s_idx in range(0, seq_len, chunk_size): + end_idx = min(s_idx + chunk_size, seq_len) + logsumexp_values[s_idx:end_idx] = torch.logsumexp(logits[s_idx:end_idx], dim=-1) + + return logsumexp_values + +def log_probs_from_logits(logits: torch.Tensor, labels: torch.Tensor, temperature: float = 1.0) -> torch.Tensor: + if temperature != 1.0: + logits.div_(temperature) + # https://github.com/OpenRLHF/OpenRLHF/pull/718#issuecomment-2641081881 + if logits.dtype in [torch.float32, torch.float64]: + batch_dim = logits.shape[:-1] + last_dim = logits.shape[-1] + try: + from flash_attn.ops.triton.cross_entropy import cross_entropy_loss + + output = cross_entropy_loss(logits.reshape(-1, last_dim), labels.reshape(-1)) + log_probs_labels = -output[0].view(*batch_dim) + except ImportError: + logits_labels = torch.gather(logits, dim=-1, index=labels.unsqueeze(-1)).squeeze(-1) + logsumexp_values = _logsumexp_by_chunk(logits.reshape(-1, last_dim)) + logsumexp_values = logsumexp_values.view(*batch_dim) + log_probs_labels = logits_labels - logsumexp_values # log_softmax(x_i) = x_i - logsumexp(x) + else: + log_probs_labels = [] + for row_logits, row_labels in zip(logits, labels): # loop to reduce peak mem consumption + row_log_probs = F.log_softmax(row_logits, dim=-1) + row_log_probs_labels = row_log_probs.gather(dim=-1, index=row_labels.unsqueeze(-1)).squeeze(-1) + log_probs_labels.append(row_log_probs_labels) + log_probs_labels = torch.stack(log_probs_labels) + return log_probs_labels + + + +import time +from contextlib import contextmanager +from collections import defaultdict + +class Profiler: + def __init__(self, log_interval: int = 10, stats_file: str = None, enable_profile: bool = False): + self.log_interval = max(1, int(log_interval)) + self.stats_file = stats_file + self.stats = defaultdict(lambda: {"count": 0, "total": 0.0, "max": 0.0}) + self._inited = False + self.enable_profile = enable_profile + + def _init_once(self): + if self._inited: + return + with open(self.stats_file, "a", encoding="utf-8") as f: + f.write("ts\tname\tcount\ttotal_s\tavg_s\tmax_s\n") + self._inited = True + + def _record(self, name: str, elapsed: float): + s = self.stats[name] + s["count"] += 1 + s["total"] += elapsed + s["max"] = max(s["max"], elapsed) + if s["count"] % self.log_interval == 0: + avg = s["total"] / s["count"] + with open(self.stats_file, "a", encoding="utf-8") as f: + f.write(f"{time.time():.3f}\t{name}\t{s['count']}\t{s['total']:.6f}\t{avg:.6f}\t{s['max']:.6f}\n") + + @contextmanager + def block(self, name: str, rank: int = 0): + if not self.enable_profile or rank != 0: + yield None + return + self._init_once() + t0 = time.perf_counter() + try: + yield None + finally: + self._record(name, time.perf_counter() - t0) \ No newline at end of file diff --git a/zoo/jericho/priorzero/src/vllm_utils/vl_engine.py b/zoo/jericho/priorzero/src/vllm_utils/vl_engine.py new file mode 100644 index 000000000..b8f6290b3 --- /dev/null +++ b/zoo/jericho/priorzero/src/vllm_utils/vl_engine.py @@ -0,0 +1,256 @@ +""" +vLLM-based VL Engine for multimodal inference. + +This module provides a vLLM wrapper for Vision-Language (VL) models, +similar to the text-only vLLM engine but with multimodal support. +""" +import vllm +from typing import List, Union, Optional, Dict, Any +from PIL import Image +import numpy as np +from loguru import logger +from transformers import AutoProcessor + + +class VLActor: + """ + vLLM Actor for Vision-Language (VL) models. + + Similar to LLMActor but with multimodal support. + Applies ChatML formatting required by Instruct-tuned models. + """ + + def __init__( + self, + model: str = None, + limit_mm_per_prompt: Optional[Dict[str, int]] = None, + **kwargs + ): + """ + Args: + model: Path to VL model + limit_mm_per_prompt: Multimodal limits (e.g., {"image": 1}) + **kwargs: Additional vLLM arguments + """ + self.kwargs = kwargs + self.limit_mm_per_prompt = limit_mm_per_prompt or {"image": 1} + self.model_path = model + + logger.info(f"Initializing VLActor with model: {model}") + logger.info(f" Multimodal limits: {self.limit_mm_per_prompt}") + + self.llm = vllm.LLM( + model=model, + limit_mm_per_prompt=self.limit_mm_per_prompt, + **self.kwargs + ) + + # Load processor/tokenizer for chat template + try: + self.processor = AutoProcessor.from_pretrained(model, trust_remote_code=True) + logger.info(f" ✓ Loaded processor for chat template") + except Exception as e: + logger.warning(f" Failed to load processor: {e}. Will use raw prompts (may cause garbled output).") + self.processor = None + + def _apply_chat_template(self, prompt: str, system_prompt: Optional[str] = None, num_images: int = 1) -> str: + """ + Apply ChatML template to convert raw user prompt into model-expected format. + + For Qwen2.5-VL / Qwen3-VL Instruct models, the expected format is: + <|im_start|>system\nYou are a helpful assistant.<|im_end|> + <|im_start|>user\n\n<|im_end|> + <|im_start|>assistant\n + + Supports multiple images by inserting multiple {"type": "image"} entries. + """ + if self.processor is None: + return prompt + + messages = [] + + if system_prompt: + messages.append({"role": "system", "content": system_prompt}) + + content = [] + for _ in range(num_images): + content.append({"type": "image"}) + content.append({"type": "text", "text": prompt}) + messages.append({"role": "user", "content": content}) + + try: + formatted = self.processor.apply_chat_template( + messages, + tokenize=False, + add_generation_prompt=True, + ) + return formatted + except Exception as e: + logger.warning(f"Failed to apply chat template: {e}. Using raw prompt.") + return prompt + + def sleep(self, level=1): + """Put the engine to sleep to free GPU memory.""" + if hasattr(self.llm, 'sleep'): + self.llm.sleep(level=level) + + def wake_up(self): + """Wake up the engine from sleep mode.""" + if hasattr(self.llm, 'wake_up'): + self.llm.wake_up() + + def update_weight(self, name, dtype, shape, weight, empty_cache=False): + """Sync a single parameter from the DeepSpeed policy model to the vLLM engine.""" + return self.llm.collective_rpc("update_weight", args=(name, dtype, shape, weight, empty_cache)) + + def update_weight_cuda_ipc(self, name, dtype, shape, ipc_handles, empty_cache=False): + return self.llm.collective_rpc("update_weight_cuda_ipc", args=(name, dtype, shape, ipc_handles, empty_cache)) + + def reset_prefix_cache(self): + """Reset prefix cache after weight update.""" + self.llm.llm_engine.reset_prefix_cache() + + def generate( + self, + images: List[Union[Image.Image, np.ndarray, List[Image.Image]]], + prompts: Optional[List[str]] = None, + prompt_token_ids: Optional[List[List[int]]] = None, + sampling_params: Any = None, + system_prompt: Optional[str] = None, + ) -> List[Any]: + """ + Generate responses for multimodal inputs. + + Args: + images: List of images or image lists + prompts: List of text prompts (mutually exclusive with prompt_token_ids) + prompt_token_ids: List of token ID lists (mutually exclusive with prompts) + sampling_params: vLLM SamplingParams + system_prompt: Optional system prompt + + Returns: + List of vLLM RequestOutput objects + """ + if prompts is None and prompt_token_ids is None: + raise ValueError("Either prompts or prompt_token_ids must be provided") + if prompts is not None and prompt_token_ids is not None: + raise ValueError("Cannot provide both prompts and prompt_token_ids") + + # Prepare multimodal inputs + inputs = [] + + if prompts is not None: + # Text prompt mode (original) + for image, prompt in zip(images, prompts): + img_list = self._normalize_images(image) + formatted_prompt = self._apply_chat_template(prompt, system_prompt=system_prompt, num_images=len(img_list)) + img_data = img_list if len(img_list) > 1 else img_list[0] + + inputs.append({ + "prompt": formatted_prompt, + "multi_modal_data": {"image": img_data}, + }) + else: + # Token IDs mode (for logprob extraction) + for image, token_ids in zip(images, prompt_token_ids): + img_list = self._normalize_images(image) + img_data = img_list if len(img_list) > 1 else img_list[0] + + inputs.append({ + "prompt_token_ids": token_ids, + "multi_modal_data": {"image": img_data}, + }) + + # Generate + responses = self.llm.generate( + inputs, + sampling_params=sampling_params, + use_tqdm=False + ) + + return responses + + def _normalize_images(self, image: Union[Image.Image, np.ndarray, List]) -> List[Image.Image]: + """Normalize image input to list of PIL Images.""" + if isinstance(image, list): + img_list = [] + for img in image: + if isinstance(img, np.ndarray): + if img.dtype != np.uint8: + img = (img * 255).astype(np.uint8) + if len(img.shape) == 3 and img.shape[0] == 3: + img = np.transpose(img, (1, 2, 0)) + img = Image.fromarray(img) + img_list.append(img) + return img_list + else: + # Single image + if isinstance(image, np.ndarray): + if image.dtype != np.uint8: + image = (image * 255).astype(np.uint8) + if len(image.shape) == 3 and image.shape[0] == 3: + image = np.transpose(image, (1, 2, 0)) + image = Image.fromarray(image) + return [image] + + +def create_vllm_vl_engine( + tensor_parallel_size: int, + pretrain: str, + max_model_len: int, + gpu_memory_utilization: float = 0.3, + vllm_enable_sleep: bool = False, + limit_mm_per_prompt: Optional[Dict[str, int]] = None, + standalone: bool = False, +): + """ + Create a vLLM engine for Vision-Language (VL) models. + + Args: + tensor_parallel_size: Number of GPUs for tensor parallelism + pretrain: Path to pretrained VL model + max_model_len: Maximum sequence length + gpu_memory_utilization: GPU memory utilization ratio + vllm_enable_sleep: Whether to enable sleep mode + limit_mm_per_prompt: Multimodal limits per prompt + standalone: If True, skip DDP-specific args (external_launcher, worker_extension_cls). + Use this for single-process evaluation scripts. + + Returns: + VLActor instance + """ + if limit_mm_per_prompt is None: + limit_mm_per_prompt = {"image": 1} + + logger.info("Creating vLLM VL engine:") + logger.info(f" Model: {pretrain}") + logger.info(f" Tensor Parallel Size: {tensor_parallel_size}") + logger.info(f" Max Model Length: {max_model_len}") + logger.info(f" GPU Memory Utilization: {gpu_memory_utilization}") + logger.info(f" Enable Sleep: {vllm_enable_sleep}") + logger.info(f" Multimodal Limits: {limit_mm_per_prompt}") + + # DDP-specific args are only needed when running under torchrun + extra_kwargs = {} + if not standalone: + extra_kwargs["worker_extension_cls"] = "vllm_utils.worker.WorkerWrap" + extra_kwargs["distributed_executor_backend"] = "external_launcher" + + vllm_engine = VLActor( + model=pretrain, + tensor_parallel_size=tensor_parallel_size, + max_model_len=max_model_len, + dtype="bfloat16", + gpu_memory_utilization=gpu_memory_utilization, + enable_sleep_mode=vllm_enable_sleep, + limit_mm_per_prompt=limit_mm_per_prompt, + trust_remote_code=True, + **extra_kwargs, + ) + + if vllm_enable_sleep: + vllm_engine.sleep() + + logger.info("✓ vLLM VL engine created successfully") + + return vllm_engine diff --git a/zoo/jericho/priorzero/src/vllm_utils/vllm_engine.py b/zoo/jericho/priorzero/src/vllm_utils/vllm_engine.py new file mode 100644 index 000000000..0908d0f6d --- /dev/null +++ b/zoo/jericho/priorzero/src/vllm_utils/vllm_engine.py @@ -0,0 +1,85 @@ +import os +import queue +from typing import Any, List +import vllm + +class LLMActor: + def __init__(self, model: str = None, **kwargs): + self.requests = {} + self.kwargs = kwargs + self.llm = vllm.LLM(model=model, **self.kwargs) + + # def update_weight(self, name, dtype, shape, empty_cache=False): + # return self.llm.collective_rpc("update_weight", args=(name, dtype, shape, empty_cache)) + + def update_weight(self, name, dtype, shape, weight, empty_cache=False): + return self.llm.collective_rpc("update_weight", args=(name, dtype, shape, weight, empty_cache)) + + def update_weight_cuda_ipc(self, name, dtype, shape, ipc_handles, empty_cache=False): + return self.llm.collective_rpc("update_weight_cuda_ipc", args=(name, dtype, shape, ipc_handles, empty_cache)) + + def reset_prefix_cache(self): + self.llm.llm_engine.reset_prefix_cache() + + def sleep(self, level=1): + self.llm.sleep(level=level) + + def wake_up(self): + self.llm.wake_up() + + def add_requests(self, sampling_params, prompt_token_ids): + """ + Process requests from rank0 and generate responses. + Since only rank0 will send requests, we don't need to track actor ranks. + """ + from vllm.inputs import TokensPrompt + self.sampling_params = sampling_params + self.requests = [TokensPrompt(prompt_token_ids=r) for r in prompt_token_ids] + + def get_responses(self): + """ + Return the responses for the actor with the given rank + """ + responses = self.llm.generate( + prompts=self.requests, + sampling_params=self.sampling_params, + use_tqdm=False + ) + self.requests = {} + return responses + + +def create_vllm_engine( + tensor_parallel_size: int, + pretrain: str, + enable_prefix_caching: bool, + max_model_len: int, + gpu_memory_utilization=None, + vllm_enable_sleep=False, +): + from packaging import version + + distributed_executor_backend = "external_launcher" + + vllm_engine = LLMActor( + model=pretrain, + worker_extension_cls="vllm_utils.worker.WorkerWrap", + tensor_parallel_size=tensor_parallel_size, + distributed_executor_backend=distributed_executor_backend, + max_model_len=max_model_len, + enable_prefix_caching=enable_prefix_caching, + dtype="bfloat16", + gpu_memory_utilization=gpu_memory_utilization, + enable_sleep_mode=vllm_enable_sleep, + ) + if vllm_enable_sleep: + vllm_engine.sleep() + return vllm_engine + + +def get_physical_gpu_id(): + import torch + + device = torch.cuda.current_device() + props = torch.cuda.get_device_properties(device) + return str(props.uuid) diff --git a/zoo/jericho/priorzero/src/vllm_utils/worker.py b/zoo/jericho/priorzero/src/vllm_utils/worker.py new file mode 100644 index 000000000..b78cad0bc --- /dev/null +++ b/zoo/jericho/priorzero/src/vllm_utils/worker.py @@ -0,0 +1,23 @@ +class WorkerWrap: + def update_weight_cuda_ipc(self, name, dtype, shape, ipc_handles=None, empty_cache=False): + import torch + from vllm_utils.vllm_engine import get_physical_gpu_id + + assert dtype == self.model_config.dtype, f"mismatch dtype: src {dtype}, dst {self.model_config.dtype}" + + handle = ipc_handles[get_physical_gpu_id()] + device_id = self.device.index + func, args = handle + list_args = list(args) + list_args[6] = device_id + weight = func(*list_args) + self.model_runner.model.load_weights(weights=[(name, weight)]) + torch.cuda.synchronize() + + def update_weight(self, name, dtype, shape, weight, empty_cache=False): # pylint: disable=R0917, W0613 + import torch + assert dtype == self.model_config.dtype, f"mismatch dtype: src {dtype}, dst {self.model_config.dtype}" + + self.model_runner.model.load_weights(weights=[(name, weight)]) + + del weight diff --git a/zoo/jericho/priorzero/vl_config.py b/zoo/jericho/priorzero/vl_config.py new file mode 100644 index 000000000..563e74983 --- /dev/null +++ b/zoo/jericho/priorzero/vl_config.py @@ -0,0 +1,688 @@ +""" +VL Configuration for PriorZero with Image Input + +This module provides configuration for using Vision-Language (VL) models +to generate action priors for image-based environments (e.g., Atari). +""" +from typing import Dict, Tuple, Optional +from easydict import EasyDict +from dataclasses import dataclass, field + + +# ============================================================================== +# Game Descriptions for VL Prompts +# ============================================================================== +GAME_DESCRIPTIONS = { + 'PongNoFrameskip-v4': ( + "This is Pong. You control the right paddle. " + "Move the paddle UP or DOWN to hit the ball past the opponent's paddle on the left. " + "Score points when the opponent misses. First to 21 points wins." + ), + 'BreakoutNoFrameskip-v4': ( + "This is Breakout. You control a paddle at the bottom of the screen. " + "Move LEFT or RIGHT to bounce the ball upward and break the colored bricks. " + "Each brick broken scores points. Don't let the ball fall below the paddle." + ), + 'SpaceInvadersNoFrameskip-v4': ( + "This is Space Invaders. You control a cannon at the bottom of the screen. " + "Move LEFT/RIGHT and FIRE to shoot the descending rows of aliens. " + "Destroy all aliens before they reach the bottom. Use shields for cover." + ), + 'QbertNoFrameskip-v4': ( + "This is Q*bert. You control Q*bert on a pyramid of cubes. " + "Jump on each cube to change its color to the target color. " + "Avoid enemies like Coily the snake. Change all cubes to complete the level." + ), + 'MsPacmanNoFrameskip-v4': ( + "This is Ms. Pac-Man. Navigate the maze eating dots and power pellets. " + "Avoid the ghosts unless you've eaten a power pellet, which lets you eat them. " + "Clear all dots to advance to the next level." + ), + 'LunarLander-v2': ( + "This is Lunar Lander. You control a spacecraft descending toward a landing pad (between two flags) at coordinate (0,0).\n" + "Goal: Land gently and perfectly horizontal on the pad.\n" + "\n" + "Rewards & Penalties:\n" + "- Closer to pad / slower speed = Positive reward.\n" + "- Tilted (not horizontal) = Continuous penalty.\n" + "- Side engine fire = -0.03 points/frame.\n" + "- Main engine fire = -0.3 points/frame (10x more expensive!).\n" + "- Crash = -100 points, Safe landing = +100 points." + ), +} + + +# ============================================================================== +# VL Model Configuration Presets +# ============================================================================== +VL_MODEL_CONFIGS = { + "Qwen2.5-VL-2b": { + "model_name": "Qwen2.5-VL", + "model_path": "/mnt/shared-storage-user/puyuan/model/Qwen2.5-VL-2B-Instruct", + "tensor_parallel_size": 1, + "gpu_memory_utilization": 0.25, + "description": "Qwen2.5-VL-2B-Instruct (smaller, faster)", + }, + "Qwen2.5-VL-3b": { + "model_name": "Qwen2.5-VL", + "model_path": "/mnt/shared-storage-user/puyuan/model/Qwen2.5-VL-3B-Instruct", + "tensor_parallel_size": 1, + "gpu_memory_utilization": 0.25, + "description": "Qwen2.5-VL-3B-Instruct", + }, + "Qwen2.5-VL-7b": { + "model_name": "Qwen2.5-VL", + "model_path": "/mnt/shared-storage-user/puyuan/model/Qwen2.5-VL-7B-Instruct", + "tensor_parallel_size": 1, + "gpu_memory_utilization": 0.35, + "description": "Qwen2.5-VL-7B-Instruct (better quality)", + }, + "Qwen3-VL-2b": { + "model_name": "Qwen3-VL", + "model_path": "/mnt/shared-storage-user/puyuan/model/Qwen3-VL-2B-Instruct", + "tensor_parallel_size": 1, + "gpu_memory_utilization": 0.25, + "description": "Qwen3-VL-2B-Instruct (smaller, faster)", + }, + "Qwen3-VL-8b": { + "model_name": "Qwen3-VL", + "model_path": "/mnt/shared-storage-user/puyuan/model/Qwen3-VL-8B-Instruct", + "tensor_parallel_size": 1, + "gpu_memory_utilization": 0.25, + "description": "Qwen3-VL-8B-Instruct", + }, +} + + +def get_available_vl_models(): + """Get list of available VL model configurations""" + return list(VL_MODEL_CONFIGS.keys()) + + +def get_vl_model_config(model_key: str) -> Dict: + """Get VL model configuration by key""" + if model_key not in VL_MODEL_CONFIGS: + available = ", ".join(get_available_vl_models()) + raise ValueError( + f"Unknown VL model key: {model_key}\n" + f"Available models: {available}" + ) + return VL_MODEL_CONFIGS[model_key] + + +def print_available_vl_models(): + """Print all available VL model configurations""" + print("\n" + "="*80) + print("Available VL Model Configurations:") + print("="*80) + for key, config in VL_MODEL_CONFIGS.items(): + print(f"\n {key}:") + print(f" Path: {config['model_path']}") + print(f" Tensor Parallel Size: {config['tensor_parallel_size']}") + print(f" GPU Memory Utilization: {config['gpu_memory_utilization']}") + print(f" Description: {config['description']}") + print("="*80 + "\n") + + +@dataclass +class PriorZeroVLConfig: + """Configuration for VL-based PriorZero (image input)""" + + # VL model settings + model_name_or_path: str = "Qwen2.5-VL-7b" + + vl_model_type: str = "qwen-vl" # 'qwen-vl', 'llava', 'internvl' + + # Game description for prompts + game_description: str = "" + + # Training settings (similar to LLM config) + enable_sft: bool = False + enable_rft: bool = True + rft_loss_weight: float = 1.0 + + # VL inference settings + temperature: float = 1.0 + max_new_tokens: int = 128 # CoT reasoning + action selection, no need for 256 + tensor_parallel_size: int = 1 + gpu_memory_utilization: float = 0.3 + + # vLLM engines + enable_vllm: bool = True + enable_prefix_caching: bool = True + use_cuda_ipc: bool = False + vllm_sync_backend: str = "nccl" # vLLM 同步参数使用的后端 + vllm_sync_with_ray: bool = False # 是否使用 ray 来同步 vLLM 参数 + vllm_tensor_parallel_size: int = 1 # 每个vllm engine使用几张GPU张量并行 (Fixed: 1.5B model should use 1 GPU) + + vllm_enable_sleep: bool = True # 是否可以休眠 + enable_vllm_is_correction: bool = False + vllm_is_truncated_threshold: Tuple[float, float] = (0.5, 5.0) + use_mispo: bool = False + mispo_token_truncated_threshold: Tuple[float, float] = (0.5, 2.0) + mispo_traj_truncated_threshold: Tuple[float, float] = (0.8, 1.2) + top_p: float = 0.95 + seed: int = 0 + reduction: str = "mean" + + # User prompt settings + user_prompt_dict: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "history_with_reward": True, + "observation_with_valid_actions": False, + })) + + + + # Prior generation settings + use_prior: bool = True # Whether to use VL prior + llm_prior_temperature: float = 2.0 # Temperature for prior distribution (aligned with LLM converged config) + + # MCTS root logits configuration + mcts_root_logits_dict: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "mode": "llm_plus_wm_logits", + # "mode": "llm_logits", + "plus_method": "fixed", + "wm_weight": 0.5, + "llm_max_weight": 0.7, + "llm_min_weight": 0.3, + "max_envsteps": 1e5, + })) + + # Evaluation settings + eval_dict: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "world_model": True, + "world_model_llm_prior": True, + "llm_prior": True, + "wm_eval_freq": 1000, + "llm_eval_freq": 100, + "eval_freq": int(20000), + })) + + attn_implementation: str = "flash_attention_2" + use_cot: bool = True + cot_weight: float = 0.1 # 控制 cot前缀token的权重,由于重点是action:,所以前缀的token权重调低 + prompt_max_len: int = 8192 # Image + prompt tokens; + generate_max_len: int = 512 # CoT + action output + bf16: bool = True + + history_length: int = 3 # Number of recent steps to include in context + + # VLM image mode: controls how many images are sent to the VL model + # "current_only": only the current frame (default, backward compatible) + # "first_and_current": first history frame + current frame (2 images max) + # "all_history": all history frames + current frame (history_length+1 images max) + vlm_image_mode: str = "current_only" + + # Prompt style: "concise" (shorter, better for small VLMs) or "legacy" (verbose, original) + prompt_style: str = "legacy" + + # Training settings + colocate_all_models: bool = True + policy_model_num_gpus: int = 1 + reference_model_num_gpus: int = 1 + deepspeed_enable_sleep: bool = True + + zero_stage: int = 2 + gradient_checkpointing: bool = False + gradient_checkpointing_use_reentrant: bool = False + + # Training mode (full or lora) + train_mode_dict: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "mode": "full", # "full" or "lora" + "lora_r": 16, + "lora_alpha": 32, + "lora_dropout": 0.05, + "lora_bias": "none", + "lora_target_modules": ( + "q_proj", + "k_proj", + "v_proj", + "o_proj", + "gate_proj", + "up_proj", + "down_proj", + ), + })) + max_norm: float = 1.0 + ds_tensor_parallel_size: int = 1 + ring_attn_size: int = 1 + + # Batch sizes + train_batch_size: int = 128 + micro_train_batch_size: int = 4 + broadcast_every: int = 1 + + # Optimizer settings + learning_rate: float = 1e-6 + adam_betas: Tuple[float, float] = (0.9, 0.95) + weight_decay: float = 0.01 + lr_scheduler: str = "cosine_with_min_lr" + lr_warmup_ratio: float = 0.03 + max_steps: int = int(1e4) + + # Loss settings + policy_loss_type: str = "ppo" + reward_func: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "format_reward": True, + "format_param": EasyDict( + {"format_weight": 0.5, } + ), + })) + advantage_type: str = "advantage_batch_norm" + eps_clip_low_high: Tuple[float, float] = (0.2, 0.2) + rft_kl_coef: float = 0.01 + entropy_loss_coef: float = 0.0 + kl_estimator: str = "k3" + + # Training schedule + train_vl_after_wm_warm_step: int = int(1e2) + vl_save_freq: int = 500 + save_path: str = "" + + # Alternating training schedule (matches LLM config) + train_schedule: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "alternate": True, + "wm_update_iters": 1e3, + "llm_update_iters": 1e2, + "start_phase": "wm", + "wm_warmup_updates": 0, + })) + + enable_world_model: bool = True + enable_rft: bool = True + max_rollout_staleness: int = 1 + # vl_fixed: If True, VL policy model is frozen (inference only, no PPO training). + # NOTE: vl_fixed=True is mutually exclusive with enable_rft=True (see validate()). + vl_fixed: bool = False + + # Value normalization + value_norm_cfg: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + 'enable_stability_optimizer': True, + 'value_norm_init_momentum': 0.9, + 'value_norm_final_momentum': 0.99, + 'value_norm_warmup_steps': 100, + 'value_norm_clip_percentile': 0.95, + 'value_norm_clip_method': "soft", + "value_norm_history_size": 1000, + })) + + def validate(self) -> None: + """ + Validate configuration consistency and raise on illegal combinations. + + Rules: + 1. enable_rft=True requires vl_fixed=False (PPO training needs a trainable policy model). + 2. When vl_fixed=True the VL inference engine is frozen AND PPO training is + disabled — the pipeline is collect-only + WM training. + 3. train_schedule.alternate=True requires enable_world_model=True. + """ + valid_image_modes = ("current_only", "first_and_current", "all_history") + if self.vlm_image_mode not in valid_image_modes: + raise ValueError( + f"[PriorZeroVLConfig] Invalid vlm_image_mode='{self.vlm_image_mode}'.\n" + f"Must be one of: {valid_image_modes}" + ) + + if self.enable_rft and self.vl_fixed: + raise ValueError( + "[PriorZeroVLConfig] Illegal config: enable_rft=True AND vl_fixed=True.\n" + " enable_rft=True → PPO training is enabled, which requires a trainable policy model.\n" + " vl_fixed=True → the VL policy model is frozen (no gradient update).\n" + "These two flags are mutually exclusive. Either:\n" + " (a) Set vl_fixed=False to enable VL PPO training, or\n" + " (b) Set enable_rft=False to run WM-only training with a frozen VL prior." + ) + + if self.train_schedule.get("alternate", False) and not self.enable_world_model: + raise ValueError( + "[PriorZeroVLConfig] Illegal config: train_schedule.alternate=True but enable_world_model=False.\n" + "Alternating schedule requires the World Model training phase." + ) + + if not self.enable_rft and not self.enable_world_model: + raise ValueError( + "[PriorZeroVLConfig] Illegal config: both enable_rft=False and enable_world_model=False.\n" + "At least one training objective must be enabled." + ) + + +def get_priorzero_vl_config( + env_id: str = 'PongNoFrameskip-v4', + seed: int = 0, + exp_name: str = None, + vl_model_key: Optional[str] = None, + use_prior: bool = True, + multi_gpu: bool = False, + quick_test: bool = False, +) -> Tuple[EasyDict, EasyDict, PriorZeroVLConfig]: + """ + Generate complete PriorZero configuration with VL for image input. + + Args: + env_id: Atari environment ID + seed: Random seed + exp_name: Experiment name + vl_model_key: VL model key (e.g., 'qwen-vl-chat', 'llava-1.5-7b') + use_prior: Whether to use VL prior + multi_gpu: Whether to use multi-GPU training + quick_test: Whether to use quick test configuration + + Returns: + main_config: Main configuration dictionary + create_config: Creation configuration + vl_config: VL configuration + """ + from zoo.atari.config.atari_env_action_space_map import atari_env_action_space_map + + # Detect environment type + is_lunarlander = 'LunarLander' in env_id + + if is_lunarlander: + action_space_size = 4 + else: + action_space_size = atari_env_action_space_map[env_id] + + # Base configuration parameters + if quick_test: + collector_env_num = 2 + num_segments = 2 + game_segment_length = 200 + evaluator_env_num = 2 + num_simulations = 5 + collect_num_simulations = 5 + eval_num_simulations = 5 + batch_size = 8 + num_layers = 1 + replay_ratio = 0.1 + else: + # collector_env_num = 8 + # num_segments = 8 + collector_env_num = 4 + num_segments = 4 + game_segment_length = 200 + evaluator_env_num = 3 + num_simulations = 25 + # collect_num_simulations = 25 + collect_num_simulations = 50 + eval_num_simulations = 25 + # eval_num_simulations = 50 + + batch_size = 256 + # num_layers = 4 + num_layers = 2 + replay_ratio = 0.25 + + num_unroll_steps = 10 + infer_context_length = 4 + + # Episode step limits + if is_lunarlander: + collect_max_episode_steps = int(10000) + eval_max_episode_steps = int(10000) + else: + collect_max_episode_steps = int(5e3) + eval_max_episode_steps = int(5e3) + + # Environment configuration + env_config = dict( + stop_value=int(1e6), + env_id=env_id, + # observation_shape=(3, 64, 64), + observation_shape=(3, 96, 96), + image_size=96, + gray_scale=False, + collector_env_num=collector_env_num, + evaluator_env_num=evaluator_env_num, + n_evaluator_episode=evaluator_env_num, + manager=dict(shared_memory=False,), + collect_max_episode_steps=collect_max_episode_steps, + eval_max_episode_steps=eval_max_episode_steps, + ) + + # Policy configuration + policy_config = dict( + type='priorzero', + multi_gpu=multi_gpu, + use_wandb=False, + learn=dict( + learner=dict( + hook=dict(save_ckpt_after_iter=1000000,), + ), + ), + model=dict( + # observation_shape=(3, 64, 64), + observation_shape=(3, 96, 96), + + action_space_size=action_space_size, + # ====== [FIX] support range must cover LunarLander reward/value range (-200 ~ +300) ====== + reward_support_range=(-300., 301., 1.), + value_support_range=(-300., 301., 1.), + norm_type="BN", + num_res_blocks=1, + num_channels=64, + world_model_cfg=dict( + norm_type="BN", + # final_norm_option_in_obs_head='LayerNorm', + # final_norm_option_in_encoder='LayerNorm', + # predict_latent_loss_type='mse', + final_norm_option_in_encoder='SimNorm', + final_norm_option_in_obs_head='SimNorm', + predict_latent_loss_type='group_kl', + support_size=601, + policy_entropy_weight=5e-3, + continuous_action_space=False, + max_blocks=num_unroll_steps, + max_tokens=2 * num_unroll_steps, + context_length=2 * infer_context_length, + device='cuda', + action_space_size=action_space_size, + num_layers=num_layers, + # num_heads=8, + # embed_dim=768, + num_heads=4, + embed_dim=256, + obs_type='image', # KEY: Image input with VL prior + env_num=max(collector_env_num, evaluator_env_num), + num_simulations=num_simulations, + game_segment_length=game_segment_length, + encoder_type='resnet', + # use_priority=True, + use_priority=False, + use_normal_head=True, + use_softmoe_head=False, + use_moe_head=False, + # optim_type='AdamW_mix_lr_wdecay', + optim_type='AdamW', + + decode_loss_mode=None, + latent_recon_loss_weight=0, + task_embed_option=None, + moe_in_transformer=False, + multiplication_moe_in_transformer=False, + ) + ), + # ====== [FIX] optimizer: AdamW -> AdamW_mix_lr_wdecay (layered lr/wd for encoder/transformer/head) ====== + # optim_type='AdamW_mix_lr_wdecay', + optim_type='AdamW', + + # weight_decay=1e-2, + learning_rate=1e-4, + num_unroll_steps=num_unroll_steps, + update_per_collect=None, + replay_ratio=replay_ratio, + batch_size=batch_size, + num_simulations=num_simulations, + td_steps=5, + train_start_after_envsteps=0, + game_segment_length=game_segment_length, + replay_buffer_size=int(5e5), + eval_freq=int(2e4), + collector_env_num=collector_env_num, + evaluator_env_num=evaluator_env_num, + + num_segments=collector_env_num, + action_type="varied_action_space", + model_path=None, + reanalyze_ratio=0, + cos_lr_scheduler=False, + fixed_temperature_value=0.25, + manual_temperature_decay=False, + n_episode=collector_env_num, + buffer_reanalyze_freq=1 / 1000000, + reanalyze_batch_size=160, + reanalyze_partition=0.75, + device='cuda', + + collect_num_simulations=collect_num_simulations, + eval_num_simulations=eval_num_simulations, + off_policy_degree=0, + enable_async_eval=False, + + # ====== [FIX] grad clip: 10 -> 5, prevent gradient explosion ====== + grad_clip_value=5, + value_loss_weight=0.25, + policy_loss_weight=1.0, + reward_loss_weight=1.0, + + # ====== [FIX] Adaptive entropy weight ====== + # use_adaptive_entropy_weight=True, + use_adaptive_entropy_weight=False, + adaptive_entropy_alpha_lr=1e-4, + target_entropy_start_ratio=0.98, + target_entropy_end_ratio=0.7, + target_entropy_decay_steps=100000, + # ====== [FIX] Encoder-clip annealing (prevents latent state norm from diverging) ====== + # use_encoder_clip_annealing=True, + use_encoder_clip_annealing=False, + encoder_clip_anneal_type='cosine', + encoder_clip_start_value=30.0, + encoder_clip_end_value=10.0, + encoder_clip_anneal_steps=100000, + # ====== [FIX] Priority Experience Replay ====== + # use_priority=True, + use_priority=False, + priority_prob_alpha=1, + priority_prob_beta=1, + # ====== [FIX] Label smoothing ====== + # policy_ls_eps_start=0.05, + policy_ls_eps_start=0.0, + policy_ls_eps_end=0.01, + policy_ls_eps_decay_steps=50000, + label_smoothing_eps=0.1, + # ====== Monitor ====== + monitor_norm_freq=10000, + ) + + main_config = EasyDict(dict( + env=env_config, + policy=policy_config, + exp_name=exp_name or f'data_priorzero_vl/{env_id}_seed{seed}', + seed=seed + )) + + if is_lunarlander: + env_create_cfg = dict( + type='lunarlander_image', + import_names=['zoo.box2d.lunarlander.envs.lunarlander_image_env'], + ) + else: + env_create_cfg = dict( + type='atari_lightzero', + import_names=['zoo.atari.envs.atari_lightzero_env'], + ) + + create_config = EasyDict(dict( + env=env_create_cfg, + env_manager=dict(type='subprocess'), + policy=dict( + type='priorzero', + import_names=['zoo.jericho.priorzero.src.priorzero_policy'], + ), + collector=dict( + type='priorzero_segment', + import_names=['zoo.jericho.priorzero.priorzero_collector_unified'], + ), + evaluator=dict( + type='priorzero', + import_names=['zoo.jericho.priorzero.src.priorzero_evaluator'], + ), + replay_buffer=dict( + type='game_buffer_muzero', + import_names=['lzero.mcts.buffer.game_buffer_muzero'], + ), + )) + + # VL configuration + vl_config = PriorZeroVLConfig(use_prior=use_prior) + + # Set game description + vl_config.game_description = GAME_DESCRIPTIONS.get(env_id, "") + + # Auto-configure VL model + if use_prior: + if vl_model_key is None: + vl_model_key = "qwen-vl-chat" # Default VL + print(f"[Config] Using default VL model: {vl_model_key}") + + vl_model_config = get_vl_model_config(vl_model_key) + vl_config.model_name_or_path = vl_model_config["model_path"] + vl_config.vl_model_type = vl_model_config["model_name"] + vl_config.tensor_parallel_size = vl_model_config["tensor_parallel_size"] + vl_config.gpu_memory_utilization = vl_model_config["gpu_memory_utilization"] + + print(f"[Config] VL configuration applied:") + print(f" - Model: {vl_model_key}") + print(f" - Path: {vl_config.model_name_or_path}") + print(f" - Tensor Parallel Size: {vl_config.tensor_parallel_size}") + print(f" - GPU Memory Utilization: {vl_config.gpu_memory_utilization}") + + # Override VL config for quick_test to avoid stuck training + if quick_test: + # Reduce WM warmup so VL training phase can be reached sooner + vl_config.train_schedule = EasyDict({ + "alternate": True, + "wm_update_iters": 50, # Reduced from 1000 + "llm_update_iters": 20, # Reduced from 100 + "start_phase": "wm", + "wm_warmup_updates": 0, + }) + # Only run WM+VL prior eval in quick_test (skip slow pure-VL and pure-WM eval) + vl_config.eval_dict = EasyDict({ + "world_model": False, + "world_model_llm_prior": True, + "llm_prior": False, + "eval_freq": int(50), # Reduced from 500 + }) + else: + print(f"[Config] VL prior disabled (use_prior=False)") + vl_config = None + + return main_config, create_config, vl_config + + +if __name__ == "__main__": + # Test configuration generation + print("PriorZero VL Configuration") + print("=" * 80) + + # List available models + print_available_vl_models() + + # Generate test config + print("\nGenerating test configuration...") + main_cfg, create_cfg, vl_cfg = get_priorzero_vl_config( + env_id='PongNoFrameskip-v4', + seed=0, + vl_model_key='qwen-vl-chat', + use_prior=True, + quick_test=True, + ) + + print("\n✓ Configuration generated successfully") + print(f" - Experiment: {main_cfg.exp_name}") + print(f" - Environment: {main_cfg.env.env_id}") + print(f" - Observation shape: {main_cfg.policy.model.observation_shape}") + print(f" - obs_type: {main_cfg.policy.model.world_model_cfg.obs_type}") + if vl_cfg: + print(f" - VL model: {vl_cfg.model_name_or_path}") + print(f" - Use prior: {vl_cfg.use_prior}") diff --git a/zoo/jericho/priorzero/vl_engine.py b/zoo/jericho/priorzero/vl_engine.py new file mode 100644 index 000000000..ce2076e0e --- /dev/null +++ b/zoo/jericho/priorzero/vl_engine.py @@ -0,0 +1,699 @@ +""" +Vision-Language (VL) Engine + +This module provides a unified interface for various VL models +to generate action priors from image observations. + +Supported models: +- Qwen-VL / Qwen2-VL / Qwen2.5-VL / Qwen3-VL (via vLLM) +- LLaVA-1.5 / LLaVA-1.6 +- InternVL +""" +import os +from typing import List, Union, Optional, Dict, Any +from pathlib import Path +from PIL import Image +import numpy as np +import torch +from loguru import logger + +try: + from vllm import LLM + VLLM_AVAILABLE = True +except ImportError: + VLLM_AVAILABLE = False + logger.warning("vLLM not available. VL engine will use transformers backend.") + + +class VLEngine: + """ + Base VL Engine class. + + Provides a unified interface for different VL implementations. + """ + + def __init__( + self, + model_name: str, + model_path: str, + device: str = "cuda", + tensor_parallel_size: int = 1, + gpu_memory_utilization: float = 0.3, + **kwargs + ): + """ + Args: + model_name: Model identifier (e.g., 'qwen-vl', 'llava-1.5') + model_path: Path to model weights + device: Device to run on + tensor_parallel_size: Number of GPUs for tensor parallelism + gpu_memory_utilization: GPU memory utilization ratio + """ + self.model_name = model_name + self.model_path = model_path + self.device = device + self.tensor_parallel_size = tensor_parallel_size + self.gpu_memory_utilization = gpu_memory_utilization + + self.model = None + self.tokenizer = None + self.processor = None + + logger.info(f"Initializing VL Engine: {model_name}") + self._load_model() + + def _load_model(self): + """Load the VL model. To be implemented by subclasses.""" + raise NotImplementedError("Subclasses must implement _load_model()") + + def generate( + self, + image: Union[Image.Image, np.ndarray], + prompt: str, + temperature: float = 1.0, + max_new_tokens: int = 512, + **kwargs + ) -> str: + """ + Generate text response from image and prompt. + + Args: + image: Input image (PIL Image or numpy array) + prompt: Text prompt + temperature: Sampling temperature + max_new_tokens: Maximum number of tokens to generate + + Returns: + Generated text response + """ + raise NotImplementedError("Subclasses must implement generate()") + + def batch_generate( + self, + images: List[Union[Image.Image, np.ndarray]], + prompts: List[str], + temperature: float = 1.0, + max_new_tokens: int = 512, + **kwargs + ) -> List[str]: + """ + Batch generate text responses. + + Args: + images: List of input images + prompts: List of text prompts + temperature: Sampling temperature + max_new_tokens: Maximum number of tokens to generate + + Returns: + List of generated text responses + """ + # Default implementation: sequential generation + results = [] + for image, prompt in zip(images, prompts): + result = self.generate(image, prompt, temperature, max_new_tokens, **kwargs) + results.append(result) + return results + + def wake_up(self): + """ + Wake up the engine (for vLLM sleep mode compatibility). + Subclasses using vLLM should override this. + """ + if hasattr(self, 'model') and hasattr(self.model, 'wake_up'): + self.model.wake_up() + # For non-vLLM engines, this is a no-op + + def sleep(self, level: int = 1): + """ + Put the engine to sleep (for vLLM sleep mode compatibility). + Subclasses using vLLM should override this. + + Args: + level: Sleep level (1 = light sleep, 2 = deep sleep) + """ + if hasattr(self, 'model') and hasattr(self.model, 'sleep'): + self.model.sleep(level=level) + # For non-vLLM engines, this is a no-op + + +class VLLMVLEngine(VLEngine): + """ + vLLM-based VL Engine for multimodal models. + + This engine uses vLLM's native multimodal support for efficient inference + with sleep/wake_up functionality for memory management. + + Supports: + - Qwen2.5-VL-2B-Instruct / Qwen2.5-VL-7B-Instruct + - Qwen3-VL-2B-Instruct + - Any vLLM-supported multimodal model + """ + + def __init__( + self, + model_name: str, + model_path: str, + device: str = "cuda", + tensor_parallel_size: int = 1, + gpu_memory_utilization: float = 0.3, + max_model_len: int = 8192, + enable_sleep: bool = True, + limit_mm_per_prompt: Optional[Dict[str, int]] = None, + standalone: bool = False, + **kwargs + ): + """ + Args: + model_name: Model identifier + model_path: Path to model weights + device: Device to run on + tensor_parallel_size: Number of GPUs for tensor parallelism + gpu_memory_utilization: GPU memory utilization ratio + max_model_len: Maximum sequence length + enable_sleep: Whether to enable sleep mode + limit_mm_per_prompt: Multimodal limits per prompt (e.g., {"image": 5}) + """ + self.max_model_len = max_model_len + self.enable_sleep = enable_sleep + self.limit_mm_per_prompt = limit_mm_per_prompt or {"image": 1} + self.standalone = standalone + + # Call parent init which will call _load_model + super().__init__( + model_name=model_name, + model_path=model_path, + device=device, + tensor_parallel_size=tensor_parallel_size, + gpu_memory_utilization=gpu_memory_utilization, + **kwargs + ) + + def _load_model(self): + """Load VL model using vLLM.""" + try: + from vllm_utils.vl_engine import create_vllm_vl_engine + + logger.info(f"Loading VL model with vLLM from {self.model_path}") + + self.model = create_vllm_vl_engine( + tensor_parallel_size=self.tensor_parallel_size, + pretrain=self.model_path, + max_model_len=self.max_model_len, + gpu_memory_utilization=self.gpu_memory_utilization, + vllm_enable_sleep=self.enable_sleep, + limit_mm_per_prompt=self.limit_mm_per_prompt, + standalone=self.standalone, + ) + + logger.info("✓ vLLM VL engine loaded successfully") + + except Exception as e: + logger.error(f"Failed to load vLLM VL engine: {e}") + raise + + def generate( + self, + image: Union[Image.Image, np.ndarray, List[Image.Image]], + prompt: str, + temperature: float = 1.0, + max_new_tokens: int = 512, + system_prompt: Optional[str] = None, + return_logprobs: bool = False, + **kwargs + ) -> Union[str, Dict[str, Any]]: + """Generate response using vLLM. Supports single image or image list.""" + from vllm import SamplingParams + + # Sampling parameters with sane defaults to prevent garbled output + sampling_params = SamplingParams( + temperature=max(temperature, 0.1), + max_tokens=max_new_tokens, + top_p=kwargs.pop('top_p', 0.95), + top_k=kwargs.pop('top_k', 50), + repetition_penalty=kwargs.pop('repetition_penalty', 1.1), + logprobs=None, + prompt_logprobs=1 if return_logprobs else None, + **kwargs + ) + + # Generate (VLActor expects lists) + outputs = self.model.generate( + images=[image], + prompts=[prompt], + sampling_params=sampling_params, + system_prompt=system_prompt, + ) + + # Extract text and logprobs from output + if outputs and len(outputs) > 0: + output = outputs[0].outputs[0] + text = output.text + if return_logprobs: + return { + 'text': text, + 'prompt_logprobs': outputs[0].prompt_logprobs if hasattr(outputs[0], 'prompt_logprobs') else None + } + return text + return "" if not return_logprobs else {'text': "", 'prompt_logprobs': None} + + def batch_generate_with_token_ids( + self, + images: List[Union[Image.Image, np.ndarray, List[Image.Image]]], + prompt_token_ids: List[List[int]], + temperature: float = 1.0, + max_new_tokens: int = 512, + return_logprobs: bool = False, + **kwargs + ) -> Union[List[str], List[Dict[str, Any]]]: + """ + Batch generate with token IDs input (aligned with LLM). + This allows extracting logprobs for the full sequence including appended actions. + """ + from vllm import SamplingParams + + sampling_params = SamplingParams( + temperature=max(temperature, 0.1), + max_tokens=max_new_tokens, + top_p=kwargs.pop('top_p', 0.95), + top_k=kwargs.pop('top_k', 50), + repetition_penalty=kwargs.pop('repetition_penalty', 1.1), + logprobs=None, + prompt_logprobs=1 if return_logprobs else None, + **kwargs + ) + + # Generate with token IDs + outputs = self.model.generate( + images=images, + prompt_token_ids=prompt_token_ids, + sampling_params=sampling_params, + ) + + # Extract results + results = [] + for output in outputs: + if output.outputs: + out = output.outputs[0] + text = out.text + if return_logprobs: + results.append({ + 'text': text, + 'prompt_logprobs': output.prompt_logprobs if hasattr(output, 'prompt_logprobs') else None + }) + else: + results.append(text) + else: + results.append("" if not return_logprobs else {'text': "", 'prompt_logprobs': None}) + + return results + + def batch_generate( + self, + images: List[Union[Image.Image, np.ndarray, List[Image.Image]]], + prompts: List[str], + temperature: float = 1.0, + max_new_tokens: int = 512, + system_prompt: Optional[str] = None, + return_logprobs: bool = False, + **kwargs + ) -> Union[List[str], List[Dict[str, Any]]]: + """Batch generate responses using vLLM. Supports single images or image lists per prompt.""" + from vllm import SamplingParams + + # Sampling parameters with sane defaults to prevent garbled output + sampling_params = SamplingParams( + temperature=max(temperature, 0.1), + max_tokens=max_new_tokens, + top_p=kwargs.pop('top_p', 0.95), + top_k=kwargs.pop('top_k', 50), + repetition_penalty=kwargs.pop('repetition_penalty', 1.1), + logprobs=None, + prompt_logprobs=1 if return_logprobs else None, + **kwargs + ) + + # Batch generate + outputs = self.model.generate( + images=images, + prompts=prompts, + sampling_params=sampling_params, + system_prompt=system_prompt, + ) + + # Extract texts and logprobs + results = [] + for output in outputs: + if output.outputs: + out = output.outputs[0] + text = out.text + if return_logprobs: + results.append({ + 'text': text, + 'prompt_logprobs': output.prompt_logprobs if hasattr(output, 'prompt_logprobs') else None + }) + else: + results.append(text) + else: + results.append("" if not return_logprobs else {'text': "", 'prompt_logprobs': None}) + + return results + + def wake_up(self): + """Wake up the vLLM engine.""" + if hasattr(self.model, 'wake_up'): + self.model.wake_up() + + def sleep(self, level: int = 1): + """Put the vLLM engine to sleep.""" + if hasattr(self.model, 'sleep'): + self.model.sleep(level=level) + + def update_weight(self, name, dtype, shape, weight, empty_cache=False): + """Sync a single parameter from DeepSpeed policy model to the vLLM VL engine.""" + return self.model.update_weight(name, dtype, shape, weight, empty_cache) + + def update_weight_cuda_ipc(self, name, dtype, shape, ipc_handles, empty_cache=False): + return self.model.update_weight_cuda_ipc(name, dtype, shape, ipc_handles, empty_cache) + + def reset_prefix_cache(self): + """Reset prefix cache after weight update.""" + self.model.reset_prefix_cache() + + +class QwenVLEngine(VLEngine): + """ + Qwen-VL / Qwen2-VL / Qwen2.5-VL / Qwen3-VL Engine + + Supports: + - Qwen-VL-Chat + - Qwen2-VL-2B-Instruct / Qwen2-VL-7B-Instruct + - Qwen2.5-VL-2B-Instruct / Qwen2.5-VL-7B-Instruct + - Qwen3-VL-2B-Instruct + """ + + def _load_model(self): + """Load Qwen-VL model.""" + try: + from transformers import AutoModelForVision2Seq, AutoTokenizer + from transformers.generation import GenerationConfig + + logger.info(f"Loading Qwen-VL from {self.model_path}") + + # Load tokenizer + self.tokenizer = AutoTokenizer.from_pretrained( + self.model_path, + trust_remote_code=True + ) + + # Load model - Use AutoModelForVision2Seq for VL models + self.model = AutoModelForVision2Seq.from_pretrained( + self.model_path, + device_map="auto" if self.tensor_parallel_size > 1 else self.device, + trust_remote_code=True, + torch_dtype=torch.bfloat16, + ).eval() + + # Set generation config + self.model.generation_config = GenerationConfig.from_pretrained( + self.model_path, + trust_remote_code=True + ) + + logger.info("✓ Qwen-VL model loaded successfully") + + except Exception as e: + logger.error(f"Failed to load Qwen-VL: {e}") + raise + + def generate( + self, + image: Union[Image.Image, np.ndarray], + prompt: str, + temperature: float = 1.0, + max_new_tokens: int = 512, + **kwargs + ) -> str: + """Generate response using Qwen-VL.""" + # Convert numpy array to PIL Image if needed + if isinstance(image, np.ndarray): + if image.dtype != np.uint8: + image = (image * 255).astype(np.uint8) + image = Image.fromarray(image) + + # Save image temporarily (Qwen-VL requires image path) + import tempfile + with tempfile.NamedTemporaryFile(suffix='.png', delete=False) as f: + image.save(f.name) + image_path = f.name + + try: + # Build query with image + query = self.tokenizer.from_list_format([ + {'image': image_path}, + {'text': prompt}, + ]) + + # Generate + response, history = self.model.chat( + self.tokenizer, + query=query, + history=None, + temperature=temperature, + max_new_tokens=max_new_tokens, + ) + + return response + + finally: + # Clean up temp file + os.unlink(image_path) + + +class LLaVAEngine(VLEngine): + """ + LLaVA Engine + + Supports: + - LLaVA-1.5-7B + - LLaVA-1.5-13B + - LLaVA-1.6-7B + """ + + def _load_model(self): + """Load LLaVA model.""" + try: + from transformers import AutoProcessor, LlavaForConditionalGeneration + + logger.info(f"Loading LLaVA from {self.model_path}") + + # Load processor and model + self.processor = AutoProcessor.from_pretrained(self.model_path) + self.model = LlavaForConditionalGeneration.from_pretrained( + self.model_path, + device_map="auto" if self.tensor_parallel_size > 1 else self.device, + torch_dtype=torch.float16, + ).eval() + + logger.info("✓ LLaVA model loaded successfully") + + except Exception as e: + logger.error(f"Failed to load LLaVA: {e}") + raise + + def generate( + self, + image: Union[Image.Image, np.ndarray], + prompt: str, + temperature: float = 1.0, + max_new_tokens: int = 512, + **kwargs + ) -> str: + """Generate response using LLaVA.""" + # Convert numpy array to PIL Image if needed + if isinstance(image, np.ndarray): + if image.dtype != np.uint8: + image = (image * 255).astype(np.uint8) + image = Image.fromarray(image) + + # Prepare inputs + conversation = [ + { + "role": "user", + "content": [ + {"type": "image"}, + {"type": "text", "text": prompt}, + ], + }, + ] + + prompt_text = self.processor.apply_chat_template( + conversation, add_generation_prompt=True + ) + + inputs = self.processor( + images=image, + text=prompt_text, + return_tensors="pt" + ).to(self.device) + + # Generate + with torch.no_grad(): + output_ids = self.model.generate( + **inputs, + max_new_tokens=max_new_tokens, + temperature=temperature, + do_sample=temperature > 0, + ) + + # Decode + response = self.processor.decode( + output_ids[0][inputs['input_ids'].shape[1]:], + skip_special_tokens=True + ) + + return response + + +class InternVLEngine(VLEngine): + """ + InternVL Engine + + Supports: + - InternVL-Chat-V1.5 + - InternVL2-2B + - InternVL2-8B + """ + + def _load_model(self): + """Load InternVL model.""" + try: + from transformers import AutoModel, AutoTokenizer + + logger.info(f"Loading InternVL from {self.model_path}") + + # Load tokenizer and model + self.tokenizer = AutoTokenizer.from_pretrained( + self.model_path, + trust_remote_code=True + ) + + self.model = AutoModel.from_pretrained( + self.model_path, + device_map="auto" if self.tensor_parallel_size > 1 else self.device, + trust_remote_code=True, + torch_dtype=torch.bfloat16, + ).eval() + + logger.info("✓ InternVL model loaded successfully") + + except Exception as e: + logger.error(f"Failed to load InternVL: {e}") + raise + + def generate( + self, + image: Union[Image.Image, np.ndarray], + prompt: str, + temperature: float = 1.0, + max_new_tokens: int = 512, + **kwargs + ) -> str: + """Generate response using InternVL.""" + # Convert numpy array to PIL Image if needed + if isinstance(image, np.ndarray): + if image.dtype != np.uint8: + image = (image * 255).astype(np.uint8) + image = Image.fromarray(image) + + # Generate + response = self.model.chat( + self.tokenizer, + pixel_values=None, + question=prompt, + generation_config={ + 'max_new_tokens': max_new_tokens, + 'temperature': temperature, + 'do_sample': temperature > 0, + }, + image=image, + ) + + return response + + +# VL Model Registry +VL_MODEL_REGISTRY = { + 'qwen-vl': QwenVLEngine, + 'qwen2-vl': QwenVLEngine, + 'qwen2.5-vl': VLLMVLEngine, # Use vLLM for Qwen2.5-VL + 'qwen3-vl': VLLMVLEngine, # Use vLLM for Qwen3-VL + 'llava': LLaVAEngine, + 'llava-1.5': LLaVAEngine, + 'llava-1.6': LLaVAEngine, + 'internvl': InternVLEngine, + 'internvl2': InternVLEngine, +} + + +def create_vl_engine( + model_name: str, + model_path: str, + device: str = "cuda", + tensor_parallel_size: int = 1, + gpu_memory_utilization: float = 0.3, + **kwargs +) -> VLEngine: + """ + Factory function to create VL engine. + + Args: + model_name: Model identifier (e.g., 'qwen-vl', 'llava-1.5') + model_path: Path to model weights + device: Device to run on + tensor_parallel_size: Number of GPUs for tensor parallelism + gpu_memory_utilization: GPU memory utilization ratio + + Returns: + VLEngine instance + """ + # Normalize model name + model_name_lower = model_name.lower() + + # Find matching engine class + engine_class = None + for key, cls in VL_MODEL_REGISTRY.items(): + if key in model_name_lower: + engine_class = cls + break + + if engine_class is None: + raise ValueError( + f"Unknown VL model: {model_name}. " + f"Supported models: {list(VL_MODEL_REGISTRY.keys())}" + ) + + # Create engine + engine = engine_class( + model_name=model_name, + model_path=model_path, + device=device, + tensor_parallel_size=tensor_parallel_size, + gpu_memory_utilization=gpu_memory_utilization, + **kwargs + ) + + return engine + + +if __name__ == "__main__": + # Example usage + print("VL Engine Module") + print("=" * 80) + print("\nSupported VL models:") + for model_name in VL_MODEL_REGISTRY.keys(): + print(f" - {model_name}") + + print("\nUsage:") + print(" engine = create_vl_engine('qwen-vl', '/path/to/model')") + print(" response = engine.generate(image, prompt)") diff --git a/zoo/textcraft/__init__.py b/zoo/textcraft/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/zoo/textcraft/priorzero/README.md b/zoo/textcraft/priorzero/README.md new file mode 100644 index 000000000..e6bd82276 --- /dev/null +++ b/zoo/textcraft/priorzero/README.md @@ -0,0 +1,37 @@ +# PriorZero TextCraft + +PriorZero (MCTS + World Model + LLM Prior) adapted for the AgentGym-RL TextCraft environment — a Minecraft-style crafting task where the agent must gather ingredients and follow recipes to craft a target item. + +## Prerequisites + +Start the TextCraft server: +```bash +cd /path/to/AgentGym-RL/AgentGym/agentenv-textcraft +python -m agentenv_textcraft.launch --port 36005 +``` + +## Quick Start + +```bash +# Full training (4 GPUs) +bash scripts/run_priorzero_ddp.sh + +# Debug mode (single GPU) +CUDA_VISIBLE_DEVICES=0 python src/priorzero_entry_sync_ddp.py \ + --env_id textcraft --env_addr http://127.0.0.1:36005 \ + --data_idx 0 --model qwen2.5-3b --use_cot --quick_test +``` + +## Configuration + +Key parameters in `scripts/run_priorzero_ddp.sh`: +- `DATA_IDX`: Selects goal item from crafting tree (sorted by depth) +- `LLM_MODEL`: Model size (`qwen2.5-0.5b`, `qwen2.5-1.5b`, `qwen2.5-3b`, `qwen2.5-7b`) +- `USE_COT`: Enable chain-of-thought reasoning (recommended: `true`) +- `AGENTGYM_SERVER_ADDR`: TextCraft server address (default: `http://127.0.0.1:36005`) + +## Environment Details + +- **Reward**: Binary (0 = not done, 1 = goal item crafted) +- **Actions**: Free-form text — `craft using `, `get `, `inventory` +- **Max steps**: 30 (aligned with AgentGym-RL baseline) diff --git a/zoo/textcraft/priorzero/__init__.py b/zoo/textcraft/priorzero/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/zoo/textcraft/priorzero/envs/__init__.py b/zoo/textcraft/priorzero/envs/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/zoo/textcraft/priorzero/envs/textcraft_env.py b/zoo/textcraft/priorzero/envs/textcraft_env.py new file mode 100644 index 000000000..ee0f29307 --- /dev/null +++ b/zoo/textcraft/priorzero/envs/textcraft_env.py @@ -0,0 +1,367 @@ +import copy +import json +import logging +import re +import time +from collections import OrderedDict +from typing import Any, Dict, List, Optional, Union + +import gym +import numpy as np +import torch +import requests +from requests.adapters import HTTPAdapter +from urllib3.util.retry import Retry +from transformers import AutoTokenizer + +from ding.utils import ENV_REGISTRY, get_rank, get_world_size +from ding.envs import BaseEnv, BaseEnvTimestep + + +class TextCraftHttpClient: + """HTTP client for AgentGym TextCraft server with retry and timeout.""" + + def __init__(self, env_addr: str, timeout: float = 10.0, max_retries: int = 3): + self._addr = env_addr.rstrip('/') + self._timeout = timeout + self._session = requests.Session() + retries = Retry( + total=max_retries, + backoff_factor=0.5, + status_forcelist=[500, 502, 503, 504], + ) + self._session.mount('http://', HTTPAdapter(max_retries=retries)) + self._session.mount('https://', HTTPAdapter(max_retries=retries)) + + def health_check(self) -> bool: + try: + r = self._session.get(f"{self._addr}/", timeout=self._timeout) + return r.status_code == 200 + except Exception: + return False + + def create(self, commands: str = None, goal: str = None) -> int: + payload = {} + if commands is not None: + payload["commands"] = commands + if goal is not None: + payload["goal"] = goal + r = self._session.post(f"{self._addr}/create", json=payload, timeout=self._timeout) + r.raise_for_status() + data = r.json() + if "error" in data: + raise RuntimeError(f"TextCraft create error: {data['error']}") + return data["id"] + + def reset(self, env_id: int, data_idx: int) -> dict: + r = self._session.post( + f"{self._addr}/reset", + json={"id": env_id, "data_idx": data_idx}, + timeout=self._timeout, + ) + r.raise_for_status() + data = r.json() + if "error" in data: + raise RuntimeError(f"TextCraft reset error: {data['error']}") + return data + + def step(self, env_id: int, action: str) -> dict: + r = self._session.post( + f"{self._addr}/step", + json={"id": env_id, "action": action}, + timeout=self._timeout, + ) + r.raise_for_status() + data = r.json() + if "error" in data: + raise RuntimeError(f"TextCraft step error: {data['error']}") + return data + + def close(self, env_id: int): + try: + self._session.post( + f"{self._addr}/close", + json={"id": env_id}, + timeout=self._timeout, + ) + except Exception: + pass + + def close_session(self): + self._session.close() + + +def _parse_goal(obs_text: str) -> str: + """Extract goal from observation. Format: '...Goal: craft .'""" + match = re.search(r'Goal:\s*craft\s+(.+?)\.?\s*$', obs_text, re.MULTILINE) + if match: + return match.group(1).strip() + return "" + + +def _extract_candidate_actions(obs_text: str) -> List[str]: + """ + Parse crafting recipes from observation text and generate candidate actions. + + Returns a list of executable commands: + - 'craft using ' for each recipe + - 'get ' for each base (non-craftable) ingredient + - 'inventory' always included + """ + candidates: List[str] = [] + craftable_items: set = set() + + recipe_section = re.search( + r'Crafting commands?:\s*\n(.*?)(?:\n\s*\n|\Z)', + obs_text, + re.DOTALL | re.IGNORECASE, + ) + if recipe_section: + recipe_block = recipe_section.group(1) + for line in recipe_block.strip().splitlines(): + line = line.strip() + if not line: + continue + cmd_match = re.match( + r'(craft\s+\d+\s+.+?\s+using\s+.+)', line, re.IGNORECASE, + ) + if cmd_match: + craft_cmd = cmd_match.group(1).strip() + candidates.append(craft_cmd) + output_match = re.match( + r'craft\s+\d+\s+(.+?)\s+using\s+', craft_cmd, re.IGNORECASE, + ) + if output_match: + craftable_items.add(output_match.group(1).strip().lower()) + + base_ingredients: OrderedDict = OrderedDict() + for cmd in candidates: + using_match = re.search(r'using\s+(.+)$', cmd, re.IGNORECASE) + if using_match: + parts = using_match.group(1).split(',') + for part in parts: + part = part.strip() + ing_match = re.match(r'(\d+)\s+(.+)', part) + if ing_match: + count = ing_match.group(1) + item = ing_match.group(2).strip() + if item.lower() not in craftable_items: + key = item.lower() + if key not in base_ingredients: + base_ingredients[key] = (count, item) + + for count, item in base_ingredients.values(): + candidates.append(f"get {count} {item}") + + candidates.append("inventory") + return candidates + + +@ENV_REGISTRY.register('textcraft') +class TextCraftEnv(BaseEnv): + """ + TextCraft environment wrapper for PriorZero. + Communicates with AgentGym TextCraft HTTP server. + Interface contract matches BabyAIEnv/JerichoEnv for algorithm-layer compatibility. + """ + tokenizer: Optional[AutoTokenizer] = None + + DEFAULT_CONFIG: Dict[str, Any] = { + 'env_addr': 'http://127.0.0.1:36005', + 'data_idx': 0, + 'max_steps': 30, + 'max_action_num': 20, + 'tokenizer_path': 'BAAI/bge-base-en-v1.5', + 'max_seq_len': 512, + 'for_unizero': True, + 'save_replay': False, + 'collector_env_num': 1, + 'evaluator_env_num': 1, + } + + def __init__(self, cfg: Dict[str, Any]) -> None: + merged_cfg = copy.deepcopy(self.DEFAULT_CONFIG) + merged_cfg.update(cfg) + self.cfg = merged_cfg + + self.env_addr: str = self.cfg['env_addr'] + self.data_idx: int = self.cfg['data_idx'] + self.max_steps: int = self.cfg['max_steps'] + self.max_action_num: int = self.cfg['max_action_num'] + self.max_seq_len: int = self.cfg['max_seq_len'] + self.for_unizero: bool = self.cfg['for_unizero'] + self.save_replay: bool = self.cfg['save_replay'] + + self.world_size: int = get_world_size() + self.rank: int = get_rank() + + if TextCraftEnv.tokenizer is None: + if self.rank == 0: + TextCraftEnv.tokenizer = AutoTokenizer.from_pretrained(self.cfg['tokenizer_path']) + if self.world_size > 1: + torch.distributed.barrier() + if self.rank != 0: + TextCraftEnv.tokenizer = AutoTokenizer.from_pretrained(self.cfg['tokenizer_path']) + + self._client = TextCraftHttpClient(self.env_addr) + try: + self._env_id: int = self._client.create() + except Exception as e: + logging.error(f"[TextCraftEnv] Failed to create env on server: {e}") + self._env_id = -1 + + self._goal: str = "" + self._action_list: List[str] = ["inventory"] + self._server_halted: bool = False + self.finished: bool = False + self._init_flag: bool = False + self.episode_return: float = 0.0 + self._timestep: int = 0 + + self.observation_space = gym.spaces.Dict() + self.action_space = gym.spaces.Discrete(self.max_action_num) + self.reward_space = gym.spaces.Box(low=-np.inf, high=np.inf, shape=(1,), dtype=np.float32) + + def prepare_obs(self, obs: str, return_str: bool = False) -> Dict[str, Any]: + raw_obs_text = obs + full_obs = obs + full_obs_str = copy.deepcopy(full_obs) + + if not return_str: + tokenized = TextCraftEnv.tokenizer( + [full_obs], truncation=True, padding="max_length", max_length=self.max_seq_len + ) + obs_attn_mask = tokenized['attention_mask'] + full_obs = np.array(tokenized['input_ids'][0], dtype=np.int32) + + action_mask = np.ones(self.max_action_num, dtype=np.int8) + + if return_str: + result = { + 'observation': full_obs, + 'action_mask': action_mask, + 'valid_actions': list(self._action_list), + 'raw_obs_text': raw_obs_text, + } + if self.for_unizero: + result['to_play'] = -1 + result['timestep'] = self._timestep + return result + else: + result = { + 'observation': full_obs, + 'obs_attn_mask': obs_attn_mask, + 'action_mask': action_mask, + 'valid_actions': list(self._action_list), + 'raw_obs_text': raw_obs_text, + } + if self.for_unizero: + result['to_play'] = -1 + result['timestep'] = self._timestep + return result + + def reset(self, return_str: bool = False) -> Dict[str, Any]: + if self._server_halted: + try: + self._env_id = self._client.create() + self._server_halted = False + except Exception: + pass + + try: + resp = self._client.reset(self._env_id, self.data_idx) + except Exception as e: + logging.warning(f"[TextCraftEnv] reset failed: {e}") + self._server_halted = True + self._goal = "" + self.finished = False + self._init_flag = True + self.episode_return = 0.0 + self._timestep = 0 + return self.prepare_obs("[Server unreachable]", return_str) + + obs_text = resp.get('observation', '') + self._goal = _parse_goal(obs_text) + self._action_list = _extract_candidate_actions(obs_text) + + self.finished = False + self._init_flag = True + self._server_halted = False + self.episode_return = 0.0 + self._timestep = 0 + + return self.prepare_obs(obs_text, return_str) + + def step(self, action: Union[int, np.ndarray, str], return_str: bool = False) -> BaseEnvTimestep: + if self._server_halted: + dummy_obs = self.prepare_obs("[Server halted]", return_str) + info = {'action_str': 'noop', 'abnormal': True, 'eval_episode_return': self.episode_return} + return BaseEnvTimestep(dummy_obs, 0.0, True, info) + + if isinstance(action, str): + action_str = action + elif isinstance(action, (int, np.integer, np.ndarray)): + action_idx = int(action.item() if isinstance(action, np.ndarray) else action) + if 0 <= action_idx < len(self._action_list): + action_str = self._action_list[action_idx] + else: + action_str = "inventory" + else: + action_str = "inventory" + + try: + resp = self._client.step(self._env_id, action_str) + except Exception as e: + logging.warning(f"[TextCraftEnv] step failed on '{action_str}': {e}") + self._server_halted = True + dummy_obs = self.prepare_obs("[Server halted]", return_str) + info = {'action_str': action_str, 'abnormal': True, 'eval_episode_return': self.episode_return, 'score': self.episode_return} + return BaseEnvTimestep(dummy_obs, 0.0, True, info) + + obs_text = resp.get('observation', '') + reward_from_server = float(resp.get('reward', 0.0)) + done = bool(resp.get('done', False)) + + self._action_list = _extract_candidate_actions(obs_text) + + step_reward = reward_from_server + self.episode_return = reward_from_server + + self._timestep += 1 + + if self._timestep >= self.max_steps: + done = True + + processed_obs = self.prepare_obs(obs_text, return_str) + info = {'action_str': action_str, 'score': self.episode_return} + + if done: + self.finished = True + info['eval_episode_return'] = self.episode_return + + return BaseEnvTimestep(processed_obs, step_reward, done, info) + + def seed(self, seed: int, dynamic_seed: bool = True) -> None: + self._seed = seed + + def close(self) -> None: + self._init_flag = False + if hasattr(self, '_client') and self._client is not None: + self._client.close(self._env_id) + + def __repr__(self) -> str: + return "LightZero TextCraft Env" + + @staticmethod + def create_collector_env_cfg(cfg: Dict[str, Any]) -> List[Dict[str, Any]]: + collector_env_num = cfg.pop('collector_env_num') + cfg = copy.deepcopy(cfg) + cfg['is_collect'] = True + return [cfg for _ in range(collector_env_num)] + + @staticmethod + def create_evaluator_env_cfg(cfg: Dict[str, Any]) -> List[Dict[str, Any]]: + evaluator_env_num = cfg.pop('evaluator_env_num') + cfg = copy.deepcopy(cfg) + cfg['is_collect'] = False + return [cfg for _ in range(evaluator_env_num)] diff --git a/zoo/textcraft/priorzero/scripts/run_priorzero_ddp.sh b/zoo/textcraft/priorzero/scripts/run_priorzero_ddp.sh new file mode 100644 index 000000000..cbac9b184 --- /dev/null +++ b/zoo/textcraft/priorzero/scripts/run_priorzero_ddp.sh @@ -0,0 +1,49 @@ +#!/bin/bash +set -x + +cd /mnt/shared-storage-user/puyuan/code/LightZero/zoo/textcraft/priorzero +export PYTHONPATH=/mnt/shared-storage-user/puyuan/code/LightZero:$PYTHONPATH + +# ============================================================================ +# PREREQUISITE: Start TextCraft server FIRST +# cd /path/to/AgentGym-RL/AgentGym/agentenv-textcraft +# python -m agentenv_textcraft.launch --port 36005 +# ============================================================================ + +# 1. Training environment parameters +CUDA_DEVICES="0,1,2,3" +NPROC_PER_NODE=4 +MASTER_PORT=24555 + +# 2. TextCraft-specific parameters +AGENTGYM_SERVER_ADDR="http://127.0.0.1:36005" +DATA_IDX=0 # selects goal item from crafting tree depth list + +# 3. Model parameters +LLM_MODEL="qwen2.5-3b" # "qwen2.5-0.5b" "qwen2.5-1.5b" "qwen2.5-3b" "qwen2.5-7b" +USE_COT=true +LOG_DIR="./data_priorzero/textcraft/run_logs" +mkdir -p "${LOG_DIR}" + +CURRENT_TIME=$(date +"%Y%m%d_%H%M%S") +LOG_FILE="${LOG_DIR}/log_dataidx${DATA_IDX}_${LLM_MODEL}_${CURRENT_TIME}.txt" + +# 4. Environment variables +export CUDA_VISIBLE_DEVICES="${CUDA_DEVICES}" +export PYTHONFAULTHANDLER=1 +export TORCH_DISTRIBUTED_DEBUG=DETAIL +export NCCL_DEBUG=INFO + +# 5. Build command +CMD_ARGS="--env_id textcraft --env_addr ${AGENTGYM_SERVER_ADDR} --data_idx ${DATA_IDX} --model ${LLM_MODEL}" + +if [ "${USE_COT}" = true ]; then + CMD_ARGS="${CMD_ARGS} --use_cot" +fi + +torchrun \ + --nproc_per_node="${NPROC_PER_NODE}" \ + --master-port="${MASTER_PORT}" \ + ./src/priorzero_entry_sync_ddp.py \ + ${CMD_ARGS} \ + 2>&1 | tee "${LOG_FILE}" diff --git a/zoo/textcraft/priorzero/scripts/test_1gpu.sh b/zoo/textcraft/priorzero/scripts/test_1gpu.sh new file mode 100644 index 000000000..9441ab228 --- /dev/null +++ b/zoo/textcraft/priorzero/scripts/test_1gpu.sh @@ -0,0 +1,7 @@ +#!/bin/bash +set -x + +cd /mnt/shared-storage-user/puyuan/code/LightZero/zoo/textcraft/priorzero +export PYTHONPATH=/mnt/shared-storage-user/puyuan/code/LightZero:$PYTHONPATH + +torchrun --nproc_per_node=1 --master-port=24556 ./src/priorzero_entry_sync_ddp.py --quick_test --env_addr http://127.0.0.1:36005 --data_idx 0 --model qwen2.5-3b --use_cot diff --git a/zoo/textcraft/priorzero/src/__init__.py b/zoo/textcraft/priorzero/src/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/zoo/textcraft/priorzero/src/priorzero_config.py b/zoo/textcraft/priorzero/src/priorzero_config.py new file mode 100644 index 000000000..172276ff3 --- /dev/null +++ b/zoo/textcraft/priorzero/src/priorzero_config.py @@ -0,0 +1,403 @@ +import os +from typing import Dict, Tuple, Optional, Any +from easydict import EasyDict +import torch.distributed as dist +from dataclasses import dataclass, field + +# ============================================================================ +# Model Configuration Presets (shared with Jericho/BabyAI version) +# ============================================================================ +MODEL_CONFIGS = { + "qwen2.5-0.5b": { + "model_name_or_path": "/mnt/shared-storage-user/puyuan/xiongjyu/models/Qwen2.5-0.5B-Instruct", + "vllm_tensor_parallel_size": 1, + "gpu_memory_utilization": 0.2, + "description": "Qwen2.5-0.5B-Instruct (smallest, fastest)", + }, + "qwen2.5-1.5b": { + "model_name_or_path": "/mnt/shared-storage-user/puyuan/xiongjyu/models/Qwen2.5-1.5B-Instruct", + "vllm_tensor_parallel_size": 1, + "gpu_memory_utilization": 0.2, + "description": "Qwen2.5-1.5B-Instruct (balanced performance)", + }, + "qwen2.5-3b": { + "model_name_or_path": "/mnt/shared-storage-user/puyuan/xiongjyu/models/Qwen2.5-3B-Instruct", + "vllm_tensor_parallel_size": 1, + "gpu_memory_utilization": 0.25, + "description": "Qwen2.5-3B-Instruct (better quality)", + }, + "qwen2.5-7b": { + "model_name_or_path": "/mnt/shared-storage-user/puyuan/model/Qwen2.5-7B-Instruct", + "vllm_tensor_parallel_size": 2, + "gpu_memory_utilization": 0.35, + "description": "Qwen2.5-7B-Instruct (high quality, needs 2+ GPUs)", + }, + "qwen2.5-14b": { + "model_name_or_path": "/mnt/shared-storage-user/puyuan/model/Qwen2.5-14B-Instruct", + "vllm_tensor_parallel_size": 4, + "gpu_memory_utilization": 0.5, + "description": "Qwen2.5-14B-Instruct (best quality, needs 4+ GPUs)", + }, +} + +def get_available_models(): + return list(MODEL_CONFIGS.keys()) + +def get_model_config(model_key: str) -> Dict: + if model_key not in MODEL_CONFIGS: + available = ", ".join(get_available_models()) + raise ValueError(f"Unknown model key: {model_key}\nAvailable models: {available}") + return MODEL_CONFIGS[model_key] + + +@dataclass +class PriorZeroLLMConfig: + model_name_or_path: str = "Qwen2.5-3B-Instruct" + local_rank: int = -1 + enable_rft: bool = True + enable_world_model: bool = True + train_mode_dict: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "mode": "full", + "lora_r": 16, + "lora_alpha": 32, + "lora_dropout": 0.05, + "lora_bias": "none", + "lora_target_modules": ( + "q_proj", "k_proj", "v_proj", "o_proj", + "gate_proj", "up_proj", "down_proj", + ), + })) + + train_schedule: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "alternate": True, + "wm_update_iters": 500, + "llm_update_iters": 100, + "start_phase": "wm", + "wm_warmup_updates": 0, + })) + + llm_prior_temperature: float = 1.0 + mcts_root_logits_dict: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "mode": "llm_plus_wm_logits", + "plus_method": "fixed", + "wm_weight": 0.5, + "llm_max_weight": 0.7, + "llm_min_weight": 0.3, + "max_envsteps": 1e5, + })) + eval_dict: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "world_model": True, + "world_model_llm_prior": True, + "llm_prior": True, + "wm_eval_freq": 500, + "llm_eval_freq": 50, + })) + + attn_implementation: str = "flash_attention_2" + history_length: int = 10 + use_cot: bool = True + cot_weight: float = 0.1 + + user_prompt_dict: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "history_with_reward": True, + "observation_with_valid_actions": False, + })) + + prompt_max_len: int = 4096 + generate_max_len: int = 1024 + bf16: bool = True + + enable_vllm: bool = True + enable_prefix_caching: bool = False + use_cuda_ipc: bool = False + enable_vllm_is_correction: bool = False + vllm_is_truncated_threshold: Tuple[float, float] = (0.5, 5.0) + use_mispo: bool = False + mispo_token_truncated_threshold: Tuple[float, float] = (0.5, 2.0) + mispo_traj_truncated_threshold: Tuple[float, float] = (0.8, 1.2) + + vllm_sync_backend: str = "nccl" + vllm_tensor_parallel_size: int = 1 + gpu_memory_utilization: float = 0.3 + vllm_enable_sleep: bool = True + temperature: float = 1.0 + top_p: float = 0.95 + seed: int = 0 + reduction: str = "mean" + + deepspeed_enable_sleep: bool = True + zero_stage: int = 2 + gradient_checkpointing: bool = False + gradient_checkpointing_use_reentrant: bool = False + max_norm: float = 1.0 + ds_tensor_parallel_size: int = 1 + + train_batch_size: int = 32 + micro_train_batch_size: int = 4 + max_rollout_staleness: int = 1 + + learning_rate: float = 1e-6 + adam_betas: Tuple[float, float] = (0.9, 0.95) + weight_decay: float = 0.01 + lr_scheduler: str = "cosine_with_min_lr" + lr_warmup_ratio: float = 0.03 + max_steps: int = int(1e4) + policy_loss_type: str = "ppo" + reward_func: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + "format_reward": True, + "format_param": EasyDict({"format_weight": 0.5}), + })) + advantage_type: str = "advantage_global_batch_norm" + eps_clip_low_high: Tuple[float, float] = (0.2, 0.2) + rft_kl_coef: float = 0.001 + entropy_loss_coef: float = 0.0 + kl_estimator: str = "k3" + + llm_save_freq: int = 1000 + save_path: str = "" + + value_norm_cfg: Optional[EasyDict] = field(default_factory=lambda: EasyDict({ + 'enable_stability_optimizer': True, + 'value_norm_init_momentum': 0.9, + 'value_norm_final_momentum': 0.99, + 'value_norm_warmup_steps': 100, + 'value_norm_clip_percentile': 0.95, + 'value_norm_clip_method': "soft", + "value_norm_history_size": 1000, + })) + + +def get_priorzero_config( + env_id: str = 'textcraft', + seed: int = 0, + exp_name: str = None, + use_cot: bool = True, + model_key: Optional[str] = "qwen2.5-3b", + multi_gpu: bool = False, + env_addr: str = 'http://127.0.0.1:36005', + data_idx: int = 0, +) -> Tuple[EasyDict, EasyDict]: + + action_space_size = 20 + max_steps = 30 + wm_encoder_option = 'legacy' + wm_model_name = '/mnt/shared-storage-user/puyuan/xiongjyu/models/bge-base-en-v1.5' + + collector_env_num = 1 + evaluator_env_num = 2 + n_episode = collector_env_num + + num_unroll_steps = 10 + infer_context_length = 4 + game_segment_length = 50 + num_layers = 2 + embed_dim = 768 + replay_ratio = 0.1 + batch_size = 64 + collect_num_simulations = 50 + eval_num_simulations = 50 + replay_buffer_size = int(3e5) + + env_config = dict( + stop_value=int(1e6), + max_steps=max_steps, + observation_shape=512, + env_id=env_id, + env_addr=env_addr, + data_idx=data_idx, + for_unizero=True, + tokenizer_path=wm_model_name, + max_action_num=action_space_size, + max_seq_len=512, + collector_env_num=collector_env_num, + evaluator_env_num=evaluator_env_num, + n_evaluator_episode=evaluator_env_num, + manager=dict(shared_memory=False), + ) + policy_config = dict( + type='priorzero', + multi_gpu=multi_gpu, + use_wandb=False, + learn=dict( + learner=dict( + hook=dict(save_ckpt_after_iter=1000000), + ), + ), + model=dict( + observation_shape=512, + action_space_size=action_space_size, + encoder_option=wm_encoder_option, + encoder_url=wm_model_name, + model_type="mlp", + continuous_action_space=False, + norm_type="LN", + world_model_cfg=dict( + norm_type="LN", + final_norm_option_in_head="LayerNorm", + final_norm_option_in_encoder="LayerNorm", + predict_latent_loss_type='mse', + policy_entropy_weight=5e-2, + continuous_action_space=False, + max_blocks=num_unroll_steps, + max_tokens=2 * num_unroll_steps, + context_length=2 * infer_context_length, + device="cuda", + action_space_size=action_space_size, + num_layers=num_layers, + num_heads=24, + embed_dim=embed_dim, + obs_type="text", + env_num=max(collector_env_num, evaluator_env_num), + decode_loss_mode=None, + latent_recon_loss_weight=0, + task_embed_option=None, + moe_in_transformer=False, + multiplication_moe_in_transformer=False, + game_segment_length=game_segment_length, + ) + ), + update_per_collect=None, + num_segments=collector_env_num, + action_type="varied_action_space", + model_path=None, + num_unroll_steps=num_unroll_steps, + reanalyze_ratio=0, + replay_ratio=replay_ratio, + batch_size=batch_size, + learning_rate=3e-4, + weight_decay=1e-4, + cos_lr_scheduler=False, + fixed_temperature_value=0.25, + manual_temperature_decay=False, + n_episode=n_episode, + train_start_after_envsteps=0, + replay_buffer_size=replay_buffer_size, + eval_freq=int(3e4), + collector_env_num=collector_env_num, + evaluator_env_num=evaluator_env_num, + buffer_reanalyze_freq=1 / 1000000, + reanalyze_batch_size=160, + reanalyze_partition=0.75, + device='cuda', + collect_num_simulations=collect_num_simulations, + eval_num_simulations=eval_num_simulations, + game_segment_length=game_segment_length, + off_policy_degree=0, + enable_async_eval=False, + optim_type='AdamW', + grad_clip_value=10.0, + value_loss_weight=0.25, + policy_loss_weight=1.0, + reward_loss_weight=1.0, + use_adaptive_entropy_weight=False, + adaptive_entropy_alpha_lr=1e-4, + use_encoder_clip_annealing=False, + encoder_clip_anneal_type='cosine', + encoder_clip_start_value=30.0, + encoder_clip_end_value=10.0, + encoder_clip_anneal_steps=100000, + use_priority=False, + priority_prob_alpha=0.6, + priority_prob_beta=0.4, + ) + + llm_config = PriorZeroLLMConfig(use_cot=use_cot) + + model_config = get_model_config(model_key) + llm_config.model_name_or_path = model_config["model_name_or_path"] + llm_config.vllm_tensor_parallel_size = model_config["vllm_tensor_parallel_size"] + llm_config.gpu_memory_utilization = model_config["gpu_memory_utilization"] + + if exp_name is None: + if llm_config.enable_rft: + exp_name = ( + f"data_priorzero/textcraft/llm_rft/priorzero_dataidx{data_idx}_{model_key}_train_{llm_config.train_mode_dict.mode}/" + f"useCot_{llm_config.use_cot}_alternate_{llm_config.train_schedule.alternate}/" + f"mcts_{llm_config.mcts_root_logits_dict.mode}_staleness_{llm_config.max_rollout_staleness}_tbs_{llm_config.train_batch_size}_use_mispo_{llm_config.use_mispo}" + ) + else: + exp_name = ( + f"data_priorzero/textcraft/llm_frozen/priorzero_dataidx{data_idx}_{model_key}_" + f"train_{llm_config.train_mode_dict.mode}" + f"useCot_{llm_config.use_cot}_seed{seed}" + ) + + priorzero_config = dict( + env=env_config, + policy=policy_config, + exp_name=exp_name, + seed=seed + ) + create_config = dict( + env=dict( + type="textcraft", + import_names=["zoo.textcraft.priorzero.envs.textcraft_env"], + ), + env_manager=dict(type="base"), + policy=dict( + type="priorzero", + import_names=["zoo.jericho.priorzero.src.priorzero_policy"], + ), + collector=dict( + type="priorzero_segment", + import_names=["zoo.jericho.priorzero.src.priorzero_collector"], + ), + evaluator=dict( + type="priorzero", + import_names=["zoo.jericho.priorzero.src.priorzero_evaluator"], + ), + replay_buffer=dict( + type='game_buffer_muzero', + import_names=['lzero.mcts.buffer.game_buffer_muzero'], + ), + ) + main_config = EasyDict(priorzero_config) + create_config = EasyDict(create_config) + + print(f"[Config] TextCraft configuration applied:") + print(f" - Model: {model_key}") + print(f" - Path: {llm_config.model_name_or_path}") + print(f" - Server: {env_addr}") + print(f" - data_idx: {data_idx} (selects goal item from crafting tree)") + + return main_config, create_config, llm_config + + +def get_priorzero_debug_config( + env_id: str = 'textcraft', + seed: int = 0, + exp_name: str = None, + use_cot: bool = True, + model_key: Optional[str] = "qwen2.5-3b", + env_addr: str = 'http://127.0.0.1:36005', + data_idx: int = 0, +) -> EasyDict: + + main_config, create_config, llm_config = get_priorzero_config( + env_id=env_id, seed=seed, exp_name=exp_name, use_cot=use_cot, + model_key=model_key, env_addr=env_addr, data_idx=data_idx, + ) + max_steps = 15 + batch_size = 8 + collect_num_simulations = 2 + eval_num_simulations = 2 + num_layers = 1 + game_segment_length = 50 + + llm_config.train_batch_size = 8 + llm_config.micro_train_batch_size = 4 + llm_config.train_schedule.wm_update_iters = 2 + llm_config.train_schedule.llm_update_iters = 1 + llm_config.eval_dict.wm_eval_freq = 2 + llm_config.eval_dict.llm_eval_freq = 1 + + main_config.env.max_steps = max_steps + main_config.policy.model.world_model_cfg.num_layers = num_layers + main_config.policy.model.world_model_cfg.game_segment_length = game_segment_length + main_config.policy.batch_size = batch_size + main_config.policy.collect_num_simulations = collect_num_simulations + main_config.policy.eval_num_simulations = eval_num_simulations + main_config.policy.update_per_collect = 2 + main_config.policy.game_segment_length = game_segment_length + + return main_config, create_config, llm_config diff --git a/zoo/textcraft/priorzero/src/priorzero_datafactory.py b/zoo/textcraft/priorzero/src/priorzero_datafactory.py new file mode 100644 index 000000000..7ac84877c --- /dev/null +++ b/zoo/textcraft/priorzero/src/priorzero_datafactory.py @@ -0,0 +1,90 @@ +import importlib.util +from pathlib import Path +from typing import List, Tuple, Optional + +_jericho_df_path = str( + Path(__file__).resolve().parent.parent.parent.parent + / "jericho" / "priorzero" / "src" / "priorzero_datafactory.py" +) +_spec = importlib.util.spec_from_file_location("jericho_datafactory", _jericho_df_path) +_jericho_mod = importlib.util.module_from_spec(_spec) +_spec.loader.exec_module(_jericho_mod) +JerichoDataProcessor = _jericho_mod.DataProcessor + + +class DataProcessor(JerichoDataProcessor): + """TextCraft-specific DataProcessor with Minecraft crafting prompts.""" + + def get_system_prompt(self): + parts = [ + "You are an expert agent in a Minecraft-style crafting environment. " + "You are given crafting recipes and must craft a target item by gathering ingredients and following recipes.", + "", + "Available action types:", + '- craft using , , ...: craft an item using a provided recipe', + '- get : obtain a raw (non-craftable) ingredient', + '- inventory: check your current inventory', + "", + "Example actions:", + " get 4 glowstone dust", + " craft 1 glowstone using 4 glowstone dust", + " craft 1 sticky piston using 1 piston, 1 slime ball", + " inventory", + "", + "Rules:", + "1. Always specify quantities in craft and get commands.", + "2. You can ONLY use crafting recipes provided in the observation. Do not invent recipes.", + "3. If a recipe uses a generic ingredient (e.g. 'planks'), you may substitute a specific type (e.g. 'dark oak planks').", + "4. Plan your crafting order: gather raw materials first, then craft intermediate items, then the final goal.", + "", + "OUTPUT FORMAT:", + ] + if self.use_cot: + parts.append( + "You MUST produce exactly TWO parts in the following order:\n" + "1. Reasoning: Analyze the goal, available recipes, current inventory, " + "and determine the next optimal action.\n" + "2. Action: A single executable command (craft/get/inventory). " + "NOT a description — the exact command to run.\n" + "Strict Format Example:\n" + "Reasoning: I need glowstone dust to craft a glowstone block. Let me get 4 glowstone dust first.\n" + "Action: get 4 glowstone dust" + ) + else: + parts.append( + "Output exactly one line starting with 'Action:' followed by the exact command.\n" + "Example:\n" + "Action: get 4 glowstone dust" + ) + return "\n".join(parts) + + def get_user_prompt(self, history=None, current_obs=None, valid_actions=None): + prompt_parts = [] + user_prompt_dict = self.args.user_prompt_dict + + if history and len(history) > 0: + prompt_parts.append("=== ACTION HISTORY ===") + for i, (obs, action, reward) in enumerate(history, start=1): + prompt_parts.append(f"Step {i}:") + prompt_parts.append(f"Observation: {obs.strip()}") + prompt_parts.append(f"Action: {action.strip()}") + if user_prompt_dict.history_with_reward: + prompt_parts.append(f"Reward: {reward}") + prompt_parts.append("") + + prompt_parts.append("=== CURRENT OBSERVATION ===") + prompt_parts.append(current_obs.strip()) + + prompt_parts.append("\n=== INSTRUCTION ===") + if self.use_cot: + prompt_parts.append( + "Analyze the observation and provide your response:\n" + "Reasoning: \n" + "Action: " + ) + else: + prompt_parts.append( + "Choose the best action:\n" + "Action: " + ) + return "\n".join(prompt_parts) diff --git a/zoo/textcraft/priorzero/src/priorzero_entry_sync_ddp.py b/zoo/textcraft/priorzero/src/priorzero_entry_sync_ddp.py new file mode 100644 index 000000000..355b4b5d3 --- /dev/null +++ b/zoo/textcraft/priorzero/src/priorzero_entry_sync_ddp.py @@ -0,0 +1,338 @@ +import sys +import os +import logging +from pathlib import Path + +_jericho_src = str(Path(__file__).resolve().parent.parent.parent.parent / "jericho" / "priorzero" / "src") +_local_src = str(Path(__file__).resolve().parent) +sys.path.insert(0, _jericho_src) +sys.path.insert(0, _local_src) + +import asyncio +from functools import partial +from typing import Tuple, Optional, List + +import torch +import torch.distributed as dist +import wandb + +from ding.config import compile_config, save_config +from ding.envs import create_env_manager, get_vec_env_setting +from ding.policy import create_policy +from ding.utils import set_pkg_seed, get_rank, get_world_size +from ding.worker import create_buffer, BaseLearner +from tensorboardX import SummaryWriter +from loguru import logger +import deepspeed + +from priorzero_config import ( + get_priorzero_config, + get_priorzero_debug_config, + get_available_models, +) +from priorzero_collector import PriorZeroCollector +from priorzero_evaluator import PriorZeroEvaluator +from priorzero_policy import * +from lzero.mcts.buffer.game_buffer_priorzero import PriorZeroGameBufferOptimized +from utils import dump_dataclass_cfg_py + +from lzero.entry.utils import calculate_update_per_collect + +def prepare_unizero(rank, cfg, create_cfg, llm_cfg, seed): + cfg = compile_config(cfg, seed=seed, auto=True, create_cfg=create_cfg) + env_fn, collector_env_cfg, evaluator_env_cfg = get_vec_env_setting(cfg.env) + collector_env = create_env_manager(cfg.env.manager, [partial(env_fn, cfg=c) for c in collector_env_cfg]) + evaluator_env = create_env_manager(cfg.env.manager, [partial(env_fn, cfg=c) for c in evaluator_env_cfg]) + + collector_env.seed(seed) + evaluator_env.seed(seed, dynamic_seed=False) + + policy = create_policy(cfg.policy, enable_field=['learn', 'collect', 'eval'], exp_name=cfg.exp_name, llm_cfg=llm_cfg) + if cfg.policy.model_path is not None: + logging.info(f"[Rank {rank}] Loading pretrained model from {cfg.policy.model_path}...") + policy.learn_mode.load_state_dict(torch.load(cfg.policy.model_path, map_location=cfg.policy.device)) + logger.info(f"[Rank {rank}] Policy created") + + os.makedirs(f'./{cfg.exp_name}/log/', exist_ok=True) + tb_logger = SummaryWriter(os.path.join(f'./{cfg.exp_name}/log/', 'serial')) if get_rank() == 0 else None + logger.info(f"[Rank {rank}] TensorBoard logger: ./{cfg.exp_name}/log/") + + learner = BaseLearner( + cfg.policy.learn.learner, + policy.learn_mode, + tb_logger, + exp_name=cfg.exp_name + ) + logger.info(f"[Rank {rank}] BaseLearner created") + + replay_buffer = PriorZeroGameBufferOptimized(cfg.policy) + logger.info(f"[Rank {rank}] PriorZero replay buffer created") + + collector = PriorZeroCollector( + env=collector_env, + policy=policy.collect_mode, + llm_config=llm_cfg, + tb_logger=tb_logger, + exp_name=cfg.exp_name, + policy_config=cfg.policy, + ) + logger.info(f"[Rank {rank}] Collector created") + + evaluator = PriorZeroEvaluator( + n_evaluator_episode=cfg.env.n_evaluator_episode, + stop_value=cfg.env.stop_value, + env=evaluator_env, + policy=policy.eval_mode, + tb_logger=tb_logger, + exp_name=cfg.exp_name, + policy_config=cfg.policy, + llm_config=llm_cfg, + ) + logger.info(f"[Rank {rank}] Evaluator created") + learner.call_hook('before_run') + + return cfg, replay_buffer, tb_logger, policy, collector, evaluator, learner + +def all_gather_cmd(world_size, obj) -> List: + if world_size <= 1: + return [obj] + lst = [None] * dist.get_world_size() + dist.all_gather_object(lst, obj) + return lst + +def train_priorzero( + cfg: dict, + create_cfg: dict, + llm_cfg, + seed: int = 0, + max_train_iter: int = int(1e6), + max_env_step: Optional[int] = int(1e10), + enable_profile: bool = False +): + rank = int(os.environ.get("RANK", "0")) + print(f"DEBUG: Is dist initialized at start? {dist.is_initialized()}") + if dist.is_initialized(): + print(f"DEBUG: Backend is {dist.get_backend()}") + from strategy.deepspeed import get_strategy, torch_dist_barrier_and_cuda_sync + strategy = get_strategy(llm_cfg) + strategy.print(llm_cfg) + + strategy.setup_distributed() + world_size = getattr(strategy, "world_size", 1) + + cfg, replay_buffer, tb_logger, policy, collector, evaluator, learner = prepare_unizero( + rank=rank, cfg=cfg, create_cfg=create_cfg, llm_cfg=llm_cfg, seed=seed + ) + batch_size = cfg.policy.batch_size + logger.info(f"[Rank {rank}] World Model components initialized") + if rank == 0: + dump_dataclass_cfg_py(llm_cfg, path=f"{cfg.exp_name}/llm_cfg.py") + llm_cfg.save_path = f'./{cfg.exp_name}/llm_ckpt/' + + from utils import Profiler + prof = Profiler(log_interval=10, stats_file=f'./{cfg.exp_name}/log/profiler.txt', enable_profile=enable_profile) + + logger.info(f"[Rank {rank}] Initializing LLM Actor...") + set_pkg_seed(seed + rank, use_cuda=True) + + from models.actor import PolicyModel, ReferenceModel + if llm_cfg.rft_kl_coef > 0: + ref_model = ReferenceModel(strategy=strategy, pretrain=llm_cfg.model_name_or_path) + else: + ref_model = None + + from vllm_utils.vllm_engine import create_vllm_engine + vllm_engine = create_vllm_engine( + tensor_parallel_size=llm_cfg.vllm_tensor_parallel_size, + pretrain=llm_cfg.model_name_or_path, + enable_prefix_caching=llm_cfg.enable_prefix_caching, + max_model_len=llm_cfg.prompt_max_len + llm_cfg.generate_max_len, + gpu_memory_utilization=llm_cfg.gpu_memory_utilization, + vllm_enable_sleep=llm_cfg.vllm_enable_sleep, + ) + print(f'[Rank {rank}] Vllm engine successfully created!') + + from priorzero_datafactory import DataProcessor + data_processor = DataProcessor( + rank=rank, world_size=world_size, vllm_engine=vllm_engine, + strategy=strategy, model_path=llm_cfg.model_name_or_path, + exp_name=cfg.exp_name if rank == 0 else None, + ) + collector.data_processor = data_processor + collector.prof = prof + evaluator.data_processor = data_processor + + policy_model = PolicyModel( + strategy=strategy, pretrain=llm_cfg.model_name_or_path, + vllm_engine=vllm_engine, max_steps=llm_cfg.max_steps + ) + from priorzero_trainer import PriorZeroLLMTrainer + trainer = PriorZeroLLMTrainer( + cfg=llm_cfg, pretrain=llm_cfg.model_name_or_path, + strategy=strategy, vllm_engine=vllm_engine, + policy_model=policy_model, reference_model=ref_model, + exp_name=cfg.exp_name if rank == 0 else None, + tb_logger=tb_logger if rank == 0 else None, + llm_save_freq=llm_cfg.llm_save_freq + ) + + torch_dist_barrier_and_cuda_sync() + train_schedule = llm_cfg.train_schedule + train_alternate = train_schedule["alternate"] + current_phase = None + if train_alternate: + current_phase = train_schedule["start_phase"] + last_wm_train_iter = 0 + last_llm_train_iter = 0 + + while True: + if collector.envstep >= max_env_step or learner.train_iter >= max_train_iter: + break + + if learner.train_iter != 0 and evaluator.should_eval(wm_train_iter=learner.train_iter, llm_train_iter=policy_model.train_iter, phase=current_phase): + logger.info(f"[Evaluator][Rank {rank}: Iter {learner.train_iter}] Evaluating...") + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.wake_up() + evaluator.eval(wm_train_iter=learner.train_iter, llm_train_iter=policy_model.train_iter, phase=current_phase) + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.sleep() + + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.wake_up() + + new_data = collector.collect(train_iter=learner.train_iter, policy_kwargs={'temperature': 0.25, 'epsilon': 0.0}, phase=current_phase) + data_processor.get_llm_output_log(wm_train_iter=learner.train_iter, llm_train_iter=policy_model.train_iter) + + if llm_cfg.vllm_enable_sleep and vllm_engine is not None: + vllm_engine.sleep() + + replay_buffer.push_game_segments(new_data) + replay_buffer.remove_oldest_data_to_fit() + num_of_transitions = replay_buffer.get_num_of_transitions() + + torch_dist_barrier_and_cuda_sync() + + if llm_cfg.enable_world_model and (not train_alternate or (train_alternate and current_phase == "wm")): + if not (num_of_transitions > batch_size): + logger.warning(f'[WM Training] Data insufficient: batch_size={batch_size}, buffer={replay_buffer}. Continue collecting...') + cmd = 0 + else: + cmd = 1 + if min(all_gather_cmd(world_size=world_size, obj=cmd)) == 0: + continue + + update_per_collect = calculate_update_per_collect(cfg, new_data, world_size=world_size) + logger.info(f"[WM Training] Rank {rank} | Iter {learner.train_iter} | Updates: {update_per_collect}") + + for i in range(update_per_collect): + with prof.block("train_world_model", rank=rank): + train_data = replay_buffer.sample(batch_size, policy) + train_data.append(learner.train_iter) + log_vars = learner.train(train_data, collector.envstep) + if cfg.policy.use_priority: + replay_buffer.update_priority(train_data, log_vars[0]['value_priority_orig']) + policy.recompute_pos_emb_diff_and_clear_cache() + if llm_cfg.enable_rft and train_alternate and learner.train_iter - last_wm_train_iter >= train_schedule["wm_update_iters"]: + current_phase = "llm" + last_wm_train_iter = learner.train_iter + replay_buffer.mark_latest_transitions_consumed() + print(f"[WM Training][Rank {rank}] Switching to LLM phase at wm iter: {learner.train_iter}") + continue + + if llm_cfg.enable_rft and (not train_alternate or (train_alternate and current_phase == "llm")): + new_num_of_transitions = replay_buffer.get_num_of_transitions() - replay_buffer.last_pos_in_transition + logger.info(f"[LLM Training] Rank {rank} | Total: {num_of_transitions} | New: {new_num_of_transitions}") + + with prof.block("fetch_latest_batch", rank=rank): + priorzero_batch = replay_buffer.fetch_latest_batch(batch_size=-1, policy=policy) + torch.cuda.empty_cache() + + with prof.block("train_llm", rank=rank): + llm_need_sample_cnt = llm_cfg.train_batch_size * llm_cfg.max_rollout_staleness // world_size + flag, train_samples = data_processor.make_llm_train_samples(priorzero_batch, ddp=True, max_samples=llm_need_sample_cnt) + + if not flag: + local_llm_ready = 0 + else: + local_llm_ready = 1 + gathered_llm_ready = all_gather_cmd(world_size=world_size, obj=local_llm_ready) + + if min(gathered_llm_ready) == 0: + logger.info(f"[Rank {rank}] Skip LLM training: not all ranks ready. flags={gathered_llm_ready}") + continue + + trainer.train_batch(train_samples, collect_env_steps=collector.envstep) + replay_buffer.mark_latest_transitions_consumed() + + torch_dist_barrier_and_cuda_sync() + if llm_cfg.enable_world_model and train_alternate and trainer.global_step - last_llm_train_iter >= train_schedule["llm_update_iters"]: + current_phase = "wm" + last_llm_train_iter = trainer.global_step + data_processor.clear_statis() + print(f"[Rank {rank}] Switching to WM phase at llm iter: {trainer.global_step}") + +def main(): + import argparse + import requests as req + + parser = argparse.ArgumentParser(description='PriorZero TextCraft Training') + parser.add_argument('--env_id', type=str, default='textcraft', help='Environment ID') + parser.add_argument('--env_addr', type=str, default='http://127.0.0.1:36005', help='TextCraft server address') + parser.add_argument('--data_idx', type=int, default=0, help='Task index (selects goal item from crafting tree)') + parser.add_argument('--seed', type=int, default=0, help='Random seed') + parser.add_argument('--max_iter', type=int, default=int(1e6), help='Max training iterations') + parser.add_argument('--quick_test', action='store_true', default=False, help='Use debug config') + parser.add_argument('--model', type=str, default="qwen2.5-3b", choices=get_available_models()) + parser.add_argument('--enable_profile', action='store_true', default=False) + parser.add_argument('--use_cot', action='store_true', default=False) + args = parser.parse_args() + + rank = int(os.environ.get("RANK", "0")) + if rank == 0: + try: + r = req.get(f"{args.env_addr}/", timeout=5) + assert r.status_code == 200, f"Server returned status {r.status_code}" + print(f"[HealthCheck] TextCraft server at {args.env_addr} is ready.") + except Exception as e: + raise RuntimeError( + f"TextCraft server not reachable at {args.env_addr}: {e}\n" + f"Start it first: cd /AgentGym/agentenv-textcraft && python -m agentenv_textcraft.launch --port 36005" + ) + + model_key = args.model + print(f"\n{'='*80}") + print(f"PriorZero TextCraft Training Configuration") + print(f"{'='*80}") + print(f"Server: {args.env_addr}") + print(f"data_idx: {args.data_idx} (selects goal item from crafting tree)") + print(f"Model: {model_key}") + print(f"Seed: {args.seed}") + print(f"Quick Test: {args.quick_test}") + print(f"CoT: {args.use_cot}") + print(f"{'='*80}\n") + + if args.quick_test: + logger.info("Using debug configuration") + main_cfg, create_cfg, llm_cfg = get_priorzero_debug_config( + args.env_id, args.seed, use_cot=args.use_cot, + exp_name=f'data_priorzero/textcraft/priorzero_debug_dataidx{args.data_idx}', + model_key=model_key, env_addr=args.env_addr, + data_idx=args.data_idx, + ) + else: + main_cfg, create_cfg, llm_cfg = get_priorzero_config( + args.env_id, args.seed, use_cot=args.use_cot, + model_key=model_key, multi_gpu=True, + env_addr=args.env_addr, data_idx=args.data_idx, + ) + + train_priorzero( + main_cfg, create_cfg, llm_cfg, + seed=args.seed, max_train_iter=args.max_iter, + enable_profile=args.enable_profile, + ) + + +if __name__ == "__main__": + os.environ['TOKENIZERS_PARALLELISM'] = 'false' + main()