Commit 1fa4964c authored by YaningGao's avatar YaningGao
Browse files

minor

parents d9b66ead 20619f61
Loading
Loading
Loading
Loading
+2 −2
Original line number Diff line number Diff line
@@ -5,7 +5,7 @@ from vagen.server.serial import serialize_observation

from .env import FrozenLakeEnv
from .env_config import FrozenLakeEnvConfig
from .service_config import FrozenLakeServiceConfig
from ..base.base_service_config import BaseServiceConfig

class FrozenLakeService(BaseService):
    """
@@ -13,7 +13,7 @@ class FrozenLakeService(BaseService):
    Implements batch operations with parallel processing for efficiency.
    """
    
    def __init__(self, config:FrozenLakeServiceConfig):
    def __init__(self, config:BaseServiceConfig):
        """
        Initialize the FrozenLakeService.
        
+0 −6
Original line number Diff line number Diff line
from vagen.env.base.base_service_config import BaseServiceConfig
from dataclasses import dataclass, fields,field

@dataclass
class FrozenLakeServiceConfig(BaseServiceConfig):
    pass
 No newline at end of file
+13 −3
Original line number Diff line number Diff line
@@ -9,7 +9,7 @@ from vagen.env.svg.score import calculate_total_score, calculate_total_score_bat
from vagen.env.svg.dino import get_dino_model
from vagen.env.svg.svg_utils import process_and_rasterize_svg, is_valid_svg
from vagen.env.utils.context_utils import parse_llm_raw_response, convert_numpy_to_PIL
from PIL import Image
from .service_config import SVGServiceConfig

class SVGService(BaseService):
    """
@@ -18,7 +18,7 @@ class SVGService(BaseService):
    Integrates DINO scoring model directly within the service.
    """
    
    def __init__(self, max_workers: int = 10, model_size: str = "small"):
    def __init__(self, config: SVGServiceConfig):
        """
        Initialize the SVGService.
        
@@ -26,17 +26,27 @@ class SVGService(BaseService):
            max_workers: Maximum number of worker threads for parallel processing
            model_size: Size of the DINO model to use ("small", "base", or "large")
        """
        self.max_workers = max_workers
        self.config= config
        self.max_workers = self.config.max_workers
        self.environments = {}
        self.env_configs = {}
        self.cache = {}
        
        # Load the DINO model directly in the service
<<<<<<< HEAD
        self.model_size = model_size
=======
        # This allows all environments to share the same model instance
        self.model_size = self.config.model_size
>>>>>>> 20619f61ee3e637ecfa735949676ca4cb0a58dc9
        self.dino_model = None  # Will be loaded on first use
        
        # Store device for model inference
        self.device = "cuda" if torch.cuda.is_available() else "cpu"
<<<<<<< HEAD
=======
        logging.info(f"SVGService initialized with {self.max_workers} workers, model_size={self.model_size}, device={self.device}")
>>>>>>> 20619f61ee3e637ecfa735949676ca4cb0a58dc9
    
    def _get_dino_model(self):
        """
+1 −1
Original line number Diff line number Diff line
@@ -3,4 +3,4 @@ from dataclasses import dataclass, fields,field

@dataclass
class SVGServiceConfig(BaseServiceConfig):
    pass
 No newline at end of file
    model_size="small"
 No newline at end of file
+4 −1
Original line number Diff line number Diff line
@@ -5,6 +5,7 @@ import importlib
from typing import Dict, List, Tuple, Optional, Any, Type
from vagen.env import REGISTERED_ENV
from vagen.env.base.base_service import BaseService
from vagen.env.base.base_service_config import BaseServiceConfig
import hydra
from omegaconf import DictConfig

@@ -196,7 +197,8 @@ class BatchEnvServer:
                raise ValueError(f"No service class registered for environment type: {env_name}")
                
            service_class = REGISTERED_ENV[env_name]["service_cls"]
            self.services[env_name] = service_class()
            service_config = REGISTERED_ENV[env_name].get("service_config", BaseServiceConfig)(**self.config.get(env_name, {}))
            self.services[env_name] = service_class(service_config)
                
        return self.services[env_name]
    
@@ -452,6 +454,7 @@ def main(cfg: DictConfig):
        cfg: Configuration object from Hydra
    """
    # Create and start server with configuration
    breakpoint()
    server = BatchEnvServer(cfg)
    
    print(f"Starting Batch Environment Server on http://{cfg.server.host}:{cfg.server.port}")
Loading