Commit 60cdef4a authored by jameskrw's avatar jameskrw
Browse files

update create_dataset

parent 06cc2221
Loading
Loading
Loading
Loading
+24 −16
Original line number Diff line number Diff line
@@ -21,14 +21,14 @@ def create_dataset_from_yaml(yaml_file_path: str, force_gen=False):
        env_name: sokoban  # or frozenlake
        env_config:
            # parameters to override the default env config
        split: train  # or test
        size: 100  # number of instances
        train_size: 100  # number of instances
        test_size:100
    env2:
        env_name: frozenlake
        env_config:
            # parameters to override the default env config
        split: test
        size: 50
        train_size: 100  # number of instances
        test_size:100
    ```
    
    If the environment config class (e.g., SokobanConfig, FrozenLakeConfig) has a 
@@ -66,8 +66,8 @@ def create_dataset_from_yaml(yaml_file_path: str, force_gen=False):
        
        env_name = value.get('env_name')
        custom_env_config = value.get('env_config', {})
        split = value.get('split', 'train')
        env_size = value.get('size', 100)
        train_size,test_size = value.get('train_size', 100)+value.get('test_size', 100)
        env_size = train_size + test_size
        
        env_config = REGISTERED_ENV[env_name]["config"](**custom_env_config)
        seeds_for_env = None
@@ -76,24 +76,32 @@ def create_dataset_from_yaml(yaml_file_path: str, force_gen=False):
            print(f"Using {len(seeds_for_env)} seeds generated by {env_name} config's generate_seeds method")
        else:
            seeds_for_env = np.random.randint(0, 2**31 - 1, size=env_size).tolist()
        for seed in seeds_for_env:
        for seed in seeds_for_env[:train_size]:
            env_settings = {
                'env_name': env_name,
                'env_config': custom_env_config,
                'seed': seed
            }
            
            instance = {
                "data_source": env_name,
                "prompt": [{"role": "user", "content": ''}],
                "extra_info": {"split": split, **env_settings}
                "extra_info": {"split": "train", **env_settings}
            }
            
            if split == 'train':
            train_instances.append(instance)
            else:
        for seed in seeds_for_env[train_size:]:
            env_settings = {
                'env_name': env_name,
                'env_config': custom_env_config,
                'seed': seed
            }
            instance = {
                "data_source": env_name,
                "prompt": [{"role": "user", "content": ''}],
                "extra_info": {"split": "test", **env_settings}
            }
            test_instances.append(instance)
            
    
    def make_map_fn(split):
        def process_fn(example, idx):
            return example
@@ -127,8 +135,8 @@ if __name__ == "__main__":
            "env_config": {
                "num_boxes": 1
            },
            "split": "train",
            "size": 2
            "train_size": 2,
            "test_size": 2,
        },
        "env2": {
            "env_name": "frozenlake",
@@ -136,8 +144,8 @@ if __name__ == "__main__":
                "is_slippery": False,
                "p":0.1
            },
            "split": "test",
            "size": 2
            "train_size": 2,
            "test_size": 2,
        }
    }
    create_dataset_from_yaml(yaml_file_path, force_gen=True)