Commit 6f5fcdc7 authored by root's avatar root
Browse files

svg update

parent 88351b22
Loading
Loading
Loading
Loading
+4 −2
Original line number Diff line number Diff line
@@ -24,6 +24,8 @@ pip install "gymnasium[toy-text]"

### SVG
```
# Download dataset from huggingface

# Additional dependencies:
pip install bs4
pip install svgpathtools
pip install cairosvg
```
 No newline at end of file
+12 −18
Original line number Diff line number Diff line
@@ -10,6 +10,7 @@ import re
import json
import logging
import random
from PIL import Image
from typing import Dict, Any, Optional, Tuple
from pathlib import Path
from datasets import Dataset
@@ -32,7 +33,7 @@ class SVGEnv(BaseEnv):
        self.img_id = None
        self.gt_svg_code = None
        self.gt_image = None
        self.gen_svg_code = ""
        self.gen_svg_code = None
        self.gen_image = None
        
        # Initialize random number generator
@@ -127,6 +128,12 @@ class SVGEnv(BaseEnv):
                    'failure_reason': 'invalid_svg'
                }
                self.failure_logger.info(json.dumps(failure_info))
                
                done = True
                info["metrics"] = metrics
                self.total_reward += self.reward
                self.gen_svg_code = None
                return self._render(init_obs=False), self.reward, done, info
        else:
            # Valid SVG code - apply format reward and process it
            self.reward += self.config.format_reward
@@ -225,22 +232,6 @@ class SVGEnv(BaseEnv):
        """Return the total reward collected so far"""
        return self.total_reward
        
    def render(self, mode='text'):
        """Render the current state of the environment
        
        Args:
            mode: Rendering mode ('text' is the only supported mode)
            
        Returns:
            String representation of the current state
        """
        assert mode == 'text', "Only text mode is supported for rendering"
        
        if not self.gen_svg_code:
            return self.gt_svg_code
        else:
            return self.gen_svg_code
        
    def close(self):
        """Close the environment and clean up resources"""
        if hasattr(self, 'failure_logger'):
@@ -260,8 +251,10 @@ class SVGEnv(BaseEnv):
        # Determine which image to show
        if init_obs:
            img = self.gt_image
        else:
        elif self.gen_svg_code:
            img = self.gen_image
        else:
            img = Image.new('RGB', (256, 256), color='white')
            
        # Set up multi-modal data with the image
        img_placeholder = self.config.get("image_placeholder", "image")
@@ -347,6 +340,7 @@ if __name__ == "__main__":
        obs, reward, done, info = env.step(action)
        print(f"Reward: {reward}")
        print(f"Done: {done}")
        print(f"obs:{obs}")
        print(f"Score components: {info.get('scores', {})}")
        
        # Test with another seed to verify determinism
+0 −2
Original line number Diff line number Diff line
@@ -109,7 +109,6 @@ def load_svg_dataset(data_dir, dataset_name, split):
        try:
            from datasets import load_from_disk
            dataset = load_from_disk(local_path)
            print(f"Successfully loaded dataset from {local_path} with {len(dataset)} examples")
            return dataset
        except Exception as e:
            print(f"Error loading from simplified path: {e}")
@@ -117,7 +116,6 @@ def load_svg_dataset(data_dir, dataset_name, split):
    try:
        print(f"Downloading dataset from HuggingFace: {dataset_name}")
        dataset = load_dataset(dataset_name, split=split)
        print(f"Successfully downloaded dataset with {len(dataset)} examples")
        
        try:
            os.makedirs(os.path.dirname(local_path), exist_ok=True)
+2 −2
Original line number Diff line number Diff line
@@ -2,7 +2,7 @@ env1:
    env_name: svg
    env_config:
        split: train
    train_size: 10000  
    train_size: 1000
    test_size: 0

env2:
@@ -10,6 +10,6 @@ env2:
    env_config:
        split: test
    train_size: 0 
    test_size: 100
    test_size: 200
    
+4 −4
Original line number Diff line number Diff line
@@ -20,14 +20,14 @@ python3 -m vagen.trainer.main_ppo \
    data.val_files=data/svg-vision-debug/test.parquet \
    data.train_batch_size=16 \
    data.max_prompt_length=1024 \
    data.max_response_length=128 \
    data.max_trajectory_length=1800 \
    data.max_response_length=648 \
    data.max_trajectory_length=3600 \
    data.image_key=images \
    data.truncation=error \
    actor_rollout_ref.model.path=Qwen/Qwen2.5-VL-3B-Instruct \
    actor_rollout_ref.actor.optim.lr=1e-6 \
    actor_rollout_ref.model.use_remove_padding=True \
    actor_rollout_ref.actor.ppo_mini_batch_size=32 \
    actor_rollout_ref.actor.ppo_mini_batch_size=4 \
    actor_rollout_ref.actor.ppo_micro_batch_size_per_gpu=1 \
    actor_rollout_ref.actor.use_kl_loss=False \
    actor_rollout_ref.actor.kl_loss_coef=0.001 \
@@ -38,7 +38,7 @@ python3 -m vagen.trainer.main_ppo \
    actor_rollout_ref.rollout.log_prob_micro_batch_size_per_gpu=1 \
    actor_rollout_ref.rollout.tensor_model_parallel_size=2 \
    actor_rollout_ref.rollout.name=vllm \
    actor_rollout_ref.rollout.gpu_memory_utilization=0.3 \
    actor_rollout_ref.rollout.gpu_memory_utilization=0.5 \
    actor_rollout_ref.rollout.enable_chunked_prefill=False \
    actor_rollout_ref.rollout.enforce_eager=False \
    actor_rollout_ref.rollout.free_cache_engine=False \