Commit c17c202f authored by jameskrw's avatar jameskrw
Browse files

untested update

parent 05f08fe9
Loading
Loading
Loading
Loading
+14 −64
Changes for vagen/env/base.py: 14 added lines, 64 removed lines.
Original line number Diff line number Diff line
@@ -8,6 +8,8 @@ from PIL import Image
import numpy as np
from dataclasses import dataclass, field

IMAGE_PLACEHOLDER = "<image>"

@dataclass
class EnvConfig:
    """
@@ -132,8 +134,13 @@ class BaseInterface(ABC):
        assert isinstance(obs["text_template"], str), f"obs['text_template'] must be str, got {type(obs['text_template'])}"
        
        if "multi_modal_data" in obs:
            for key, image in obs["multi_modal_data"].items():
            if IMAGE_PLACEHOLDER in obs["multi_modal_data"]:
                assert isinstance(obs["multi_modal_data"][IMAGE_PLACEHOLDER], list), f"obs['multi_modal_data']['<image>'] must be list, got {type(obs['multi_modal_data'][IMAGE_PLACEHOLDER])}"
                for image in obs["multi_modal_data"][IMAGE_PLACEHOLDER]:
                    assert isinstance(image, Image.Image), f"image must be PIL.Image.Image, got {type(image)}"
                len_of_images = len(obs["multi_modal_data"][IMAGE_PLACEHOLDER])
                len_of_image_in_text_template = len(re.findall(r'IMAGE_PLACEHOLDER', obs["text_template"]))
                assert len_of_images == len_of_image_in_text_template, f"len_of_images must be equal to len_of_image_in_text_template, got {len_of_images} and {len_of_image_in_text_template}"
        return obs, reward, done, info
    
            
@@ -143,14 +150,17 @@ class BaseInterface(ABC):
        obs, info = self._reset(seed)
        assert isinstance(info, dict), f"info must be dict, got {type(info)}"
        assert isinstance(obs, dict), f"obs must be dict, got {type(obs)}"
        assert "llm_raw_response" in info, f"info must contain 'llm_raw_response' key"
        assert isinstance(info["llm_raw_response"], str), f"info['llm_raw_response'] must be str, got {type(info['llm_raw_response'])}"
        assert "text_template" in obs, f"obs must contain 'text_template' key"
        assert isinstance(obs["text_template"], str), f"obs['text_template'] must be str, got {type(obs['text_template'])}"
        
        if "multi_modal_data" in obs:
            for key, image in obs["multi_modal_data"].items():
            if IMAGE_PLACEHOLDER in obs["multi_modal_data"]:
                assert isinstance(obs["multi_modal_data"][IMAGE_PLACEHOLDER], list), f"obs['multi_modal_data']['<image>'] must be list, got {type(obs['multi_modal_data'][IMAGE_PLACEHOLDER])}"
                for image in obs["multi_modal_data"][IMAGE_PLACEHOLDER]:
                    assert isinstance(image, Image.Image), f"image must be PIL.Image.Image, got {type(image)}"
                len_of_images = len(obs["multi_modal_data"][IMAGE_PLACEHOLDER])
                len_of_image_in_text_template = len(re.findall(r'IMAGE_PLACEHOLDER', obs["text_template"]))
                assert len_of_images == len_of_image_in_text_template, f"len_of_images must be equal to len_of_image_in_text_template, got {len_of_images} and {len_of_image_in_text_template}"
        return obs, info
    
    def get_traj_reward(self) -> float:
@@ -158,63 +168,3 @@ class BaseInterface(ABC):
        return self.traj_reward
    




def preprocess_text(text: str) -> dict:
    """Preprocess the raw text from llm to a list of strings

    1. Extract think from the first <think> ... </think>
    2. Extract answer from the first <answer> ... </answer>
    3. Split the answer by comma into a list of strings
    
    Args:
        text: raw text from llm

    Returns:
        dict with keys: llm_raw_response, think, answer_list
    """
    # Extract content from <think> tags if they exist
    think_match = re.search(r'<think>(.*?)</think>', text, re.DOTALL)
    
    # Extract content from <answer> tags
    answer_match = re.search(r'<answer>(.*?)</answer>', text, re.DOTALL)

    answer_list, thinking, answer_content = [], "", ""
    
    if think_match:
        thinking = think_match.group(1).strip()
    
    if answer_match:
        # Get the answer content and split by comma
        answer_content = answer_match.group(1).strip()
        # Split by comma and strip whitespace from each item
        answer_list = [item.strip() for item in answer_content.split(',') if item.strip()]
    
    return {
        'llm_raw_response': text,
        'answer_list': answer_list,
        'think': thinking,
        'answer': answer_content
    }

def convert_numpy_to_PIL(numpy_array: np.ndarray) -> Image.Image:
        """Convert a numpy array to a PIL RGB image."""
        if numpy_array.shape[-1] == 3:
            # Convert numpy array to RGB PIL Image
            return Image.fromarray(numpy_array, mode='RGB')
        else:
            raise ValueError(f"Unsupported number of channels: {numpy_array.shape[-1]}. Expected 3 (RGB).")


if __name__ == "__main__":
    text = """
    <think>
    I am thinking about the problem.
    </think>
    <answer>
    answer1, answer2, answer3
    </answer>
    """
    print(preprocess_text(text))
 No newline at end of file

vagen/env/config/sokoban.yaml

deleted100644 → 0
+0 −9
Changes for vagen/env/config/sokoban.yaml: 0 added lines, 9 removed lines.
Original line number Diff line number Diff line
env:
  name: "sokoban"
  data_dir: "data/sokoban"
  env_config:
    dim_room: [6, 6]
    num_boxes: 1
    max_steps: 100
    search_depth: 30 # this will change the starting position of the player
    visual_env: true
 No newline at end of file
+5 −6
Changes for vagen/env/create_dataset.py: 5 added lines, 6 removed lines.
Original line number Diff line number Diff line
@@ -8,12 +8,11 @@ from pathlib import Path

class DatasetCreator:

    def __init__(self, config_path):
        with open(config_path, 'r') as f:
            self.config = yaml.safe_load(f)
        self.env_name = self.config['env']['name']
        self.env_config = self.config['env']['env_config']
        self.data_dir = self.config['env']['data_dir']
    def __init__(self, config):
        self.config = config
        self.env_name = self.config['name']
        self.env_config = self.config['env_config']
        self.data_dir = self.config['data_dir']
        
        

+22 −2
Changes for vagen/env/sokoban/create_dataset.py: 22 added lines, 2 removed lines.
Original line number Diff line number Diff line
@@ -34,13 +34,33 @@ class SokobanDatasetCreator(DatasetCreator):

if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument('--config_path', type=str, default='vagen/env/config/sokoban.yaml')
    parser.add_argument('--start_seed', type=int, default=0)
    parser.add_argument('--train_size', type=int, default=100)
    parser.add_argument('--test_size', type=int, default=100)
    parser.add_argument('--force-gen', action='store_true', 
                        help='Force dataset generation even if files already exist')
    # Added arguments based on the YAML config
    parser.add_argument('--dim_room', type=int, nargs=2, default=[6, 6],
                        help='Dimensions of the room [height, width]')
    parser.add_argument('--num_boxes', type=int, default=1,
                        help='Number of boxes in the environment')
    parser.add_argument('--max_steps', type=int, default=100,
                        help='Maximum number of steps allowed')
    parser.add_argument('--search_depth', type=int, default=30,
                        help='Search depth that affects the starting position of the player')
    parser.add_argument('--visual_env', action='store_true',
                        help='Whether to use visual environment')
    parser.add_argument('--data_dir', type=str, default='data/sokoban',)

    args = parser.parse_args()
    creator = SokobanDatasetCreator(config_path=args.config_path)
    args.name = 'sokoban'
    args.env_config = {
        'dim_room': args.dim_room,
        'num_boxes': args.num_boxes,
        'max_steps': args.max_steps,
        'search_depth': args.search_depth,
        'visual_env': args.visual_env
    }
    creator = SokobanDatasetCreator(config=vars(args))
    #creator.create_filtered_dataset(start_seed=args.start_seed, train_size=args.train_size, test_size=args.test_size)
    creator.create_dataset(start_seed=args.start_seed, train_size=args.train_size, test_size=args.test_size)
+10 −8
Changes for vagen/env/sokoban/env.py: 10 added lines, 8 removed lines.
Original line number Diff line number Diff line
@@ -14,10 +14,11 @@ from vagen.env.sokoban.room_utils import generate_room
from vagen.env.base import (
    BaseEnv,
    BaseInterface,
    preprocess_text,
    convert_numpy_to_PIL,
    IMAGE_PLACEHOLDER
)

from vagen.env.utils import preprocess_text, convert_numpy_to_PIL

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.
"""
@@ -47,6 +48,9 @@ Box on target: +1.0
All boxes placed: +10.0
"""




init_observation_template = """
[Initial Observation]:
{observation}
@@ -60,8 +64,6 @@ reward: {reward}
done: {done}
"""

image_placeholder = "<image{index}>"


class SokobanEnv(BaseEnv, GymSokobanEnv):

@@ -318,7 +320,7 @@ class SokobanInterface(BaseInterface):
            else:
                break

        observation = image_placeholder.format(index=1) if not isinstance(env_state, str) else env_state
        observation = IMAGE_PLACEHOLDER if not isinstance(env_state, str) else env_state
        text_template = action_template.format(
            answer=answer,
            valid_action=valid_action,
@@ -333,7 +335,7 @@ class SokobanInterface(BaseInterface):
            obs = {
                'text_template': text_template,
                'multi_modal_data': {
                    observation: env_state,
                    IMAGE_PLACEHOLDER: [env_state],
                },
            }
        return obs, reward, done, info
@@ -401,7 +403,7 @@ class SokobanInterface(BaseInterface):
        env_state = self.env._render(mode='text' if not self.visual_env else 'rgb_array') # NOTE currently called after reset
        if isinstance(env_state, np.ndarray):
            env_state = convert_numpy_to_PIL(env_state)
        observation = image_placeholder.format(index=1) if not isinstance(env_state, str) else env_state
        observation = IMAGE_PLACEHOLDER if not isinstance(env_state, str) else env_state
        text_template = init_observation_template.format(
            observation=observation,
        )
@@ -411,7 +413,7 @@ class SokobanInterface(BaseInterface):
            obs = {
                'text_template': text_template,
                'multi_modal_data': {
                    observation: env_state,
                    IMAGE_PLACEHOLDER: [env_state],
                },
            }
        return obs, {}
Loading