Commit 81270ba2 authored by williamzhangNU's avatar williamzhangNU
Browse files

move pre/^Cst precess function to utils

parent 209c176b
Loading
Loading
Loading
Loading
+1 −0
Original line number Diff line number Diff line
@@ -8,6 +8,7 @@ from PIL import Image
import numpy as np
from dataclasses import dataclass, field


IMAGE_PLACEHOLDER = "<image>"

@dataclass
+11 −110
Original line number Diff line number Diff line
@@ -5,7 +5,6 @@ import re
import copy
from typing import Tuple, Dict, Optional, List, Any, Union
from PIL import Image
from dataclasses import dataclass

from vagen.utils import NoLoggerWarnings
from vagen.utils import set_seed
@@ -17,7 +16,11 @@ from vagen.env.base import (
    IMAGE_PLACEHOLDER
)

from vagen.env.utils import preprocess_text, convert_numpy_to_PIL
from vagen.env.utils import (
    convert_numpy_to_PIL,
    preprocess,
    postprocess,
)

system_prompt = """
You are a helpful assistant. You first think about the reasoning process in the mind and then provides the user with the answer.
@@ -184,22 +187,7 @@ class SokobanEnv(BaseEnv, GymSokobanEnv):



@dataclass
class PreprocessResult:
    action_list: List[int]
    answer_list: List[str] # string of extracted answer (may be invalid action)
    think: str
    answer: str
    llm_raw_response: str

    def to_dict(self):
        return {
            'action_list': self.action_list,
            'answer_list': self.answer_list,
            'think': self.think,
            'answer': self.answer,
            'llm_raw_response': self.llm_raw_response,
        }


@register(name="sokoban")
@@ -265,97 +253,6 @@ class SokobanInterface(BaseInterface):
        
        return cls.INVALID_ACTION
    
    @classmethod
    def _preprocess(cls, text: str) -> PreprocessResult:
        """Preprocess the raw text from LLM into a list of actions.
        NOTE Only keep valid actions.

        Args:
            text: raw text from LLM

        Returns:
            PreprocessResult containing parsed valid actions
        """
        first_step_preprocess = preprocess_text(text)
        preprocess_result = PreprocessResult(
            action_list=[],
            answer_list=first_step_preprocess['answer_list'],
            think=first_step_preprocess['think'],
            answer=first_step_preprocess['answer'],
            llm_raw_response=text,
        )
        
        # for answer in preprocess_result.answer_list:
        #     action = cls._extract_one_action(answer)
        #     if action != cls.INVALID_ACTION:
        #         preprocess_result.action_list.append(action)
        #         preprocess_result.valid_list.append(True)
        #     else:
        #         preprocess_result.action_list.append(cls.INVALID_ACTION)
        #         preprocess_result.valid_list.append(False)

        # ensure there are only valid actions
        for answer in preprocess_result.answer_list:
            action = cls._extract_one_action(answer)
            if action != cls.INVALID_ACTION:
                preprocess_result.action_list.append(action)
            else:
                break
        # preprocess_result.action_list = preprocess_result.action_list[:cls.MAX_ACTION_PER_STEP]
        return preprocess_result
        
    @classmethod
    def _postprocess(
        cls, 
        env_state: Union[str, np.ndarray], 
        reward: float,
        done: bool,
        info: Dict,
        preprocess_result: PreprocessResult,
    ) -> Tuple[Dict, float, bool, Dict]:
        """Postprocess the environment feedback
        NOTE now assume there's only one image in the observation

        Args:
            env_state: environment state (text or numpy array (image))
            reward: reward of the environment
            done: whether the environment is done
            info: extra info
            preprocess_result: preprocess result

        Returns:
            Tuple[Dict, float, bool, Dict]: observation, reward, done, info
        """

        if isinstance(env_state, np.ndarray):
            env_state = convert_numpy_to_PIL(env_state)

        answer = preprocess_result.answer
        valid_action = []
        for action in preprocess_result.action_list:
            valid_action.append(cls.ACTION_LOOKUP[action])

        observation = IMAGE_PLACEHOLDER if not isinstance(env_state, str) else env_state
        text_template = action_template.format(
            answer=answer,
            valid_action=valid_action,
            observation=observation,
            reward=reward,
            done=done,
        )

        if isinstance(env_state, str):
            obs = {'text_template': text_template}
        else:
            obs = {
                'text_template': text_template,
                'multi_modal_data': {
                    IMAGE_PLACEHOLDER: [env_state],
                },
            }
        return obs, reward, done, info

    

    def _step(self, raw_text: str) -> Tuple[Any, float, bool, Dict]:
        """Step the environment with llm raw response
@@ -377,7 +274,8 @@ class SokobanInterface(BaseInterface):
        reward, done, final_info = 0, False, {}


        preprocess_result = self._preprocess(raw_text)
        # preprocess_result = self._preprocess(raw_text)
        preprocess_result = preprocess(raw_text, self._extract_one_action, self.INVALID_ACTION)
        think = preprocess_result.think
        action_list = preprocess_result.action_list
        answer = preprocess_result.answer
@@ -405,12 +303,14 @@ class SokobanInterface(BaseInterface):

        env_state = self.env._render(mode='text' if not self.visual_env else 'rgb_array') # NOTE currently called after step

        return self._postprocess(
        return postprocess(
            env_state=env_state,
            reward=reward,
            done=done,
            info=final_info,
            preprocess_result=preprocess_result,
            action_lookup=self.ACTION_LOOKUP,
            action_template=action_template,
        )
    
    def _reset(self, seed: Optional[int] = None) -> Tuple[Dict, Dict]:
@@ -472,3 +372,4 @@ class SokobanInterface(BaseInterface):
    
    def get_traj_reward(self):
        return self.traj_reward
+113 −0
Original line number Diff line number Diff line
import re
from PIL import Image
import numpy as np
from dataclasses import dataclass
from vagen.env.base import BaseEnv, IMAGE_PLACEHOLDER
from typing import List, Dict, Tuple, Union



def preprocess_text(text: str) -> dict:
@@ -49,6 +53,115 @@ def convert_numpy_to_PIL(numpy_array: np.ndarray) -> Image.Image:
            raise ValueError(f"Unsupported number of channels: {numpy_array.shape[-1]}. Expected 3 (RGB).")



@dataclass
class PreprocessResult:
    action_list: List # list of valid action defined in the action space
    answer_list: List[str] # string of extracted answer (may be invalid action)
    think: str
    answer: str
    llm_raw_response: str

    def to_dict(self):
        return {
            'action_list': self.action_list,
            'answer_list': self.answer_list,
            'think': self.think,
            'answer': self.answer,
            'llm_raw_response': self.llm_raw_response,
        }

def preprocess(text: str, extract_action_func, invalid_action_code=0) -> PreprocessResult:
    """Preprocess the raw text from LLM into a list of actions.
    
    Args:
        text: Raw text from LLM
        extract_action_func: Function to extract action from text, should return action ID or invalid_action_code
            Function signature: extract_action_func(text: str) -> int
        invalid_action_code: Code representing an invalid action (default: 0)

    Returns:
        PreprocessResult containing parsed valid actions
    """
    parsed_text = preprocess_text(text)
    
    # Process actions until first invalid action
    action_list = []
    for answer in parsed_text['answer_list']:
        action = extract_action_func(answer)
        if action == invalid_action_code:
            break
        action_list.append(action)
    
    return PreprocessResult(
        action_list=action_list,
        answer_list=parsed_text['answer_list'],
        think=parsed_text['think'],
        answer=parsed_text['answer'],
        llm_raw_response=text
    )


def postprocess(
    env_state: Union[str, np.ndarray], 
    reward: float,
    done: bool,
    info: Dict,
    preprocess_result: PreprocessResult,
    action_lookup: Dict,
    action_template: str,
) -> Tuple[Dict, float, bool, Dict]:
    """Postprocess the environment feedback to obs, reward, done, info
    NOTE now assume there's only one image in the observation

    Args:
        env_state: environment state (text or numpy array (image))
        reward: reward of the environment
        done: whether the environment is done
        info: extra info
        preprocess_result: preprocess result
        action_lookup: action lookup to convert action space to text
        text_template: text template

    Returns:
        Tuple[Dict, float, bool, Dict]: observation, reward, done, info
    """

    if isinstance(env_state, np.ndarray):
        env_state = convert_numpy_to_PIL(env_state)

    answer = preprocess_result.answer
    valid_action = []
    for action in preprocess_result.action_list:
        valid_action.append(action_lookup[action])

    observation = IMAGE_PLACEHOLDER if not isinstance(env_state, str) else env_state
    text_template = action_template.format(
        answer=answer,
        valid_action=valid_action,
        observation=observation,
        reward=reward,
        done=done,
    )

    if isinstance(env_state, str):
        obs = {'text_template': text_template}
    else:
        obs = {
            'text_template': text_template,
            'multi_modal_data': {
                IMAGE_PLACEHOLDER: [env_state],
            },
        }
    return obs, reward, done, info








if __name__ == "__main__":
    text = """
    <think>