""" PyTorch Modeling class for MrPong Ping Pong RL Agent (fromziro/MrPong). Compatible with Hugging Face AutoModel via trust_remote_code=True. """ from typing import Optional, Tuple, Union, Dict, Any from dataclasses import dataclass import numpy as np import torch import torch.nn as nn from transformers import PreTrainedModel from transformers.utils import ModelOutput try: from .configuration_mr_pong import MrPongConfig except ImportError: try: from configuration_mr_pong import MrPongConfig except ImportError: from configuration_mrpong import MrPongConfig @dataclass class MrPongOutput(ModelOutput): """ Model output for MrPong. """ logits: torch.FloatTensor = None value_estimate: Optional[torch.FloatTensor] = None action: Optional[torch.LongTensor] = None class MrPongForRL(PreTrainedModel): config_class = MrPongConfig base_model_prefix = "mr_pong" def __init__(self, config: MrPongConfig): super().__init__(config) self.config = config act_fn = nn.Tanh if config.activation == "tanh" else (nn.ReLU if config.activation == "relu" else nn.GELU) layers = [] in_dim = config.obs_dim for h_dim in config.hidden_dims: layers.append(nn.Linear(in_dim, h_dim)) layers.append(act_fn()) in_dim = h_dim self.trunk = nn.Sequential(*layers) self.actor = nn.Linear(in_dim, config.action_dim) self.critic = nn.Linear(in_dim, 1) self.post_init() def forward( self, observation: torch.FloatTensor, deterministic: bool = True, return_dict: Optional[bool] = None, **kwargs ) -> Union[Tuple[torch.FloatTensor, torch.FloatTensor], MrPongOutput]: """ Forward pass returning action logits, state-value estimates, and greedily/sampled chosen action. """ return_dict = return_dict if return_dict is not None else self.config.use_return_dict if not isinstance(observation, torch.Tensor): observation = torch.tensor(observation, dtype=torch.float32) if observation.ndim == 1: observation = observation.unsqueeze(0) # Slice or pad to expected input dimension if observation.shape[-1] > self.config.obs_dim: observation = observation[..., :self.config.obs_dim] elif observation.shape[-1] < self.config.obs_dim: pad_size = self.config.obs_dim - observation.shape[-1] observation = nn.functional.pad(observation, (0, pad_size)) features = self.trunk(observation) logits = self.actor(features) value = self.critic(features) if deterministic: action = torch.argmax(logits, dim=-1) else: dist = torch.distributions.Categorical(logits=logits) action = dist.sample() if not return_dict: return logits, value, action return MrPongOutput( logits=logits, value_estimate=value, action=action ) @torch.no_grad() def act(self, observation: Union[np.ndarray, list, torch.Tensor], deterministic: bool = True) -> int: """ High-level inference method returning single integer action (0: Stay, 1: Up, 2: Down). """ self.eval() if not isinstance(observation, torch.Tensor): observation = torch.tensor(observation, dtype=torch.float32, device=self.device) else: observation = observation.to(self.device) out = self.forward(observation, deterministic=deterministic, return_dict=True) return out.action.squeeze().item()