Commit 7aa4337e authored by williamzhangNU's avatar williamzhangNU
Browse files

strict parse

parent 7ae77e1c
Loading
Loading
Loading
Loading
+13 −2
Original line number Diff line number Diff line
@@ -14,6 +14,7 @@ import argparse
import datasets
import multiprocessing as mp
from functools import partial
from collections import defaultdict
import numpy as np

from vagen.env.create_dataset import DatasetCreator
@@ -57,6 +58,7 @@ class SokobanDatasetCreator(DatasetCreator):
        """
        train_file_path = os.path.join(self.data_dir, 'train.parquet')
        test_file_path = os.path.join(self.data_dir, 'test.parquet')
        action_count = defaultdict(int)
        
        # Check if files already exist and force_gen is False
        if not force_gen and os.path.exists(train_file_path) and os.path.exists(test_file_path):
@@ -73,12 +75,16 @@ class SokobanDatasetCreator(DatasetCreator):
        pool.close()
        pool.join()

        valid_seeds = [seed for seed, gt_action_sequence in results if gt_action_sequence and len(gt_action_sequence) <= max_action_length]
        valid_seeds_with_actions = [(seed, gt_action_sequence) for seed, gt_action_sequence in results if gt_action_sequence and len(gt_action_sequence) <= max_action_length]
        valid_seeds = [seed for seed, _ in valid_seeds_with_actions]
        train_size = int(len(valid_seeds) * train_ratio)
        test_size = len(valid_seeds) - train_size
        print(f"Train size: {train_size}, Test size: {test_size}")
        # Analyze statistics of action sequences
        action_lengths = [len(gt_action_sequence) for _, gt_action_sequence in results if gt_action_sequence]
        action_lengths = [len(gt_action_sequence) for _, gt_action_sequence in valid_seeds_with_actions]
        for _, gt_action_sequence in valid_seeds_with_actions:
            for action in gt_action_sequence:
                action_count[action] += 1
        


@@ -110,6 +116,11 @@ class SokobanDatasetCreator(DatasetCreator):
            percentage = (count / len(action_lengths)) * 100
            print(f"  Length {length}: {count} instances ({percentage:.2f}%)")
        
        print("\nAction frequency:")
        for action, count in sorted(action_count.items(), key=lambda x: x[1], reverse=True):
            percentage = (count / len(valid_seeds_with_actions)) * 100
            print(f"  {action}: {count} instances ({percentage:.2f}%)")



        self.create_dataset(valid_seeds, train_size, test_size, force_gen=force_gen)
+33 −8
Original line number Diff line number Diff line
@@ -7,20 +7,43 @@ from typing import List, Dict, Tuple, Union



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
def preprocess_text(text: str, strict_match: bool = True) -> dict:
    """Preprocess the raw text from llm to a list of strings.
    
    Args:
        text: raw text from llm
        strict_match: If True, enforces strict format where text must contain exactly one 
                     <think>...</think> followed by one <answer>...</answer> with possible 
                     whitespace between them. No content should appear before <think> or 
                     after </answer>.

    Returns:
        dict with keys: llm_raw_response, think, answer_list
    """
    # Extract content from <think> tags if they exist
    
    if strict_match:
        strict_match_result = None

        # Verify exactly one occurrence of each tag
        think_open_count = text.count("<think>")
        think_close_count = text.count("</think>")
        answer_open_count = text.count("<answer>")
        answer_close_count = text.count("</answer>")
        tags_are_balanced = (think_open_count == 1 and think_close_count == 1 and 
                            answer_open_count == 1 and answer_close_count == 1)
        
        if tags_are_balanced:
            strict_format_pattern = r'^\s*<think>(.*?)</think>\s*<answer>(.*?)</answer>\s*$'
            strict_match_result = re.match(strict_format_pattern, text, re.DOTALL)
        if strict_match_result is None:
            return {
                'llm_raw_response': text,
                'answer_list': [],
                'think': "",
                'answer': "",
            }
    
    # Extract content from <think> tags
    think_match = re.search(r'<think>(.*?)</think>', text, re.DOTALL)
    
    # Extract content from <answer> tags
@@ -41,9 +64,11 @@ def preprocess_text(text: str) -> dict:
        'llm_raw_response': text,
        'answer_list': answer_list,
        'think': thinking,
        'answer': answer_content
        '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: