Commit 12c0a7a9 authored by williamzhangNU's avatar williamzhangNU
Browse files

add interface config

parent 4762deec
Loading
Loading
Loading
Loading
+3 −1
Original line number Diff line number Diff line
@@ -18,6 +18,7 @@ class EnvConfig:
    """
    env_name: str
    env_config: Dict[str, Any]
    interface_config: Dict[str, Any]
    seed: int

class BaseEnv(ABC):
@@ -79,8 +80,9 @@ class BaseEnv(ABC):
    
        
class BaseInterface(ABC):
    def __init__(self, **env_config):
    def __init__(self, env_config: Dict, interface_config: Dict = None):
        self.env_config = env_config
        self.interface_config = interface_config
        
    @classmethod
    def name_repr(cls) -> str:
+6 −22
Original line number Diff line number Diff line
@@ -5,15 +5,17 @@ import os
import pandas as pd
import argparse
from pathlib import Path
from typing import Union, List
from typing import Union, List, Dict

class DatasetCreator:

    def __init__(self, config):
    def __init__(self, config: Dict):
        self.config = config
        self.data_dir = self.config['data_dir']

        self.env_name = self.config['name']
        self.env_config = self.config['env_config']
        self.data_dir = self.config['data_dir']
        self.interface_config = self.config['interface_config']
        
        

@@ -39,6 +41,7 @@ class DatasetCreator:
            env_settings = {
                'env_name': self.env_name,
                'env_config': self.env_config,
                'interface_config': self.interface_config,
                'seed': seed_idx
            }

@@ -122,22 +125,3 @@ class DatasetCreator:
        except Exception as e:
            print(f"Error merging parquet files: {str(e)}")
            return False
 No newline at end of file



if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    parser.add_argument('--config_path', type=str, default='config.yaml')
    parser.add_argument('--start_seed', type=int, default=10000)
    parser.add_argument('--train_size', type=int, default=10000)
    parser.add_argument('--test_size', type=int, default=1000)
    parser.add_argument('--force-gen', action='store_true', 
                        help='Force dataset generation even if files already exist')
    args = parser.parse_args()
    creator = DatasetCreator(config_path=args.config_path)
    creator.create_dataset(
        start_seed=args.start_seed, 
        train_size=args.train_size, 
        test_size=args.test_size,
        force_gen=args.force_gen
    )
 No newline at end of file
+13 −1
Original line number Diff line number Diff line
@@ -22,7 +22,7 @@ from vagen.env.sokoban.room_utils import get_shortest_action_path, plot_animatio
class SokobanDatasetCreator(DatasetCreator):

    def _process_seed(self, seed: int, max_action_length: int = 5):
        env_interface = SokobanInterface(**self.env_config)
        env_interface = SokobanInterface(self.env_config, self.interface_config)
        env_interface.reset(seed=seed)
        gt_action_sequence = get_shortest_action_path(
            env_interface.env.room_fixed, 
@@ -141,6 +141,13 @@ if __name__ == "__main__":
    parser.add_argument('--visual_env', action='store_true',
                        help='Whether to use visual environment')
    
    parser.add_argument('--max_action_per_step', type=int, default=1,
                        help='Maximum number of actions per step')
    parser.add_argument('--max_action_penalty', type=float, default=-0.1,
                        help='Penalty for exceeding the maximum number of actions per step')
    parser.add_argument('--format_reward', type=float, default=0.5,
                        help='Reward for correct formatting')
    
    import os
    if 'PYTHONHASHSEED' not in os.environ:
        os.environ['PYTHONHASHSEED'] = '0'
@@ -158,6 +165,11 @@ if __name__ == "__main__":
        'search_depth': args.search_depth,
        'visual_env': args.visual_env
    }
    args.interface_config = {
        'max_action_per_step': args.max_action_per_step,
        'max_action_penalty': args.max_action_penalty,
        'format_reward': args.format_reward,
    }
    creator = SokobanDatasetCreator(config=vars(args))
    if args.max_action_length:
        creator.create_filtered_dataset(
+46 −22
Original line number Diff line number Diff line
@@ -161,9 +161,15 @@ class SokobanInterface(BaseInterface):

    def __init__(
            self,
            **env_config,
            env_config: Dict,
            interface_config: Dict,
        ):
        super().__init__(**env_config)
        """
        Args:
            env_config (Dict): environment configuration
            interface_config (Dict): interface configuration
        """
        super().__init__(env_config)

        dim_room = self.env_config['dim_room']
        num_boxes = self.env_config['num_boxes']
@@ -177,6 +183,15 @@ class SokobanInterface(BaseInterface):
        )
        self.visual_env = self.env_config.get('visual_env', True)

        max_action_per_step = interface_config.setdefault('max_action_per_step', 1)
        max_action_penalty = interface_config.setdefault('max_action_penalty', -0.5)
        format_reward = interface_config.setdefault('format_reward', 0.5)
        self.interface_config = {
            'max_action_per_step': max_action_per_step,
            'max_action_penalty': max_action_penalty,
            'format_reward': format_reward,
        }
        
    @classmethod
    def _extract_one_action(cls, text):
        """
@@ -233,22 +248,23 @@ class SokobanInterface(BaseInterface):
        answer = preprocess_result.answer
        final_info['llm_raw_response'] = preprocess_result.llm_raw_response

        if think and answer: # format reward for <think>...</think><answer>...</answer>
            reward += self.FORMAT_REWARD
        else: # format penalty
            reward += self.FORMAT_PENALTY
        if action_list: # valid action reward
            reward += self.VALID_ACTION_REWARD
        if len(action_list) > self.MAX_ACTION_PER_STEP:
            reward += self.MAX_ACTION_PENALTY

        # parse format and action list
        if action_list:
            reward += self.interface_config['format_reward']
        if len(action_list) > self.interface_config['max_action_per_step']:
            reward += self.interface_config['max_action_penalty']
            action_list = action_list[:self.interface_config['max_action_per_step']]
            preprocess_result.action_list = action_list
            

        info = {}
        for action in action_list:
            if done or self.env.finished():
                break
            _, env_reward, done, info = self.env.step(action)
            if env_reward == -0.1:
                env_reward = 0 # NOTE hard coded here to set step reward to 0
            # if env_reward == -0.1:
            #     env_reward = -0.01 # NOTE hard coded here to set step reward to 0
            reward += env_reward
        self.traj_reward += reward
        final_info.update(info) # NOTE currently only use the last step info
@@ -293,12 +309,13 @@ class SokobanInterface(BaseInterface):
        self.env.close()

    @classmethod
    def config_repr(cls, config: Dict) -> str:
    def config_repr(cls, env_config: Dict, interface_config: Dict) -> str:
        """
        Create a string representation of the configuration.
        
        Args:
            config: Dictionary containing configuration
            env_config: Dictionary containing environment configuration
            interface_config: Dictionary containing interface configuration
            
        Returns:
            String representation of the configuration
@@ -306,19 +323,26 @@ class SokobanInterface(BaseInterface):
        Raises:
            ValueError: If required keys are missing from the configuration
        """

        required_keys = ['dim_room', 'num_boxes', 'max_steps', 'search_depth']
        
        # Check for required keys
        if not all(key in config for key in required_keys):
            missing_keys = [key for key in required_keys if key not in config]
        if not all(key in env_config for key in required_keys):
            missing_keys = [key for key in required_keys if key not in env_config]
            raise ValueError(f"Missing required keys in config: {missing_keys}")
            
        # Format the configuration string
        return (f"SokobanGame(dim_room={config['dim_room']}, "
                f"num_boxes={config['num_boxes']}, "
                f"max_steps={config['max_steps']}, "
                f"search_depth={config['search_depth']})")
    
        env_config_str = (
            f"SokobanGame(dim_room={env_config['dim_room']}, "
            f"num_boxes={env_config['num_boxes']}, "
            f"max_steps={env_config['max_steps']}, "
            f"search_depth={env_config['search_depth']})"
        )
        interface_config_str = (
            f"SokobanInterface(max_action_per_step={interface_config.get('max_action_per_step', 1)}, "
            f"max_action_penalty={interface_config.get('max_action_penalty', -0.1)}, "
            f"format_reward={interface_config.get('format_reward', 0.5)})"
        )
        return f"{env_config_str}, {interface_config_str}"
    def get_task_instruction(self) -> str:
        return instruction_template
    
+4 −1
Original line number Diff line number Diff line
@@ -108,8 +108,8 @@ def postprocess(
    done: bool,
    info: Dict,
    preprocess_result: PreprocessResult,
    action_lookup: Dict,
    action_template: str,
    action_lookup: Dict = None,
) -> 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
@@ -133,6 +133,9 @@ def postprocess(
    answer = preprocess_result.answer
    valid_action = []
    for action in preprocess_result.action_list:
        if action_lookup is None:
            valid_action.append(action)
        else:
            valid_action.append(action_lookup[action])

    observation = IMAGE_PLACEHOLDER if not isinstance(env_state, str) else env_state
Loading