You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Ray RLlib中使用Torch Save/Load方法替代Checkpoint实现模型全量保存、续训与高效推理的技术问询

Answers to Your RLlib/PPO Questions & Inference Performance Fixes

Let’s walk through each of your questions one by one, then dive into solving that memory bottleneck and slow checkpoint restore issue you’re facing.


1. Can I save a full PyTorch model when ray.tune.run training finishes?

Absolutely! RLlib’s default checkpoints store the entire agent state (including optimizer parameters, training metrics, and configs), but you can extract and save just the PyTorch model weights separately. Here are two reliable methods:

Method 1: Save via Callback During Training

Use a custom callback to save the model once training completes:

from ray.tune import Callback
import torch

class SavePyTorchModelCallback(Callback):
    def on_train_result(self, iteration, trials, trial, result, **info):
        # Only save when reaching the final training iteration
        if result["training_iteration"] == trial.config["stop"]["training_iteration"]:
            agent = trial.trainer
            policy = agent.get_policy()  # Get the PPO policy
            # Save the model's state dict
            torch.save(policy.model.state_dict(), f"{trial.logdir}/final_pytorch_model.pt")

# Add this callback to your tune.run call
results = tune.run(
    trainer,
    config=config1,
    verbose=0,
    stop=stop1,
    checkpoint_at_end=True,
    metric='episode_reward_mean',
    mode="max",
    checkpoint_freq=1,
    callbacks=[SavePyTorchModelCallback()]
)

Method 2: Save After Training Completes

Extract the best trial’s agent and save the model directly:

import torch

best_trial = results.get_best_trial(metric='episode_reward_mean', mode="max")
best_agent = best_trial.trainer
policy = best_agent.get_policy()
torch.save(policy.model.state_dict(), "./best_pytorch_model.pt")

This saves only the model weights, which is far lighter than a full RLlib checkpoint.


2. Can I import a PyTorch model instead of restoring a checkpoint for follow-up ray.tune.run training?

Yes, but note this is a warm-start scenario (you’re starting training from pre-trained model weights, not resuming exactly where you left off). RLlib checkpoints include optimizer states and training progress, which a standalone PyTorch model doesn’t. If warm-starting fits your use case, here’s how to do it:

import torch

def load_pretrained_weights(trainer):
    policy = trainer.get_policy()
    pretrained_state = torch.load("./best_pytorch_model.pt")
    policy.model.load_state_dict(pretrained_state)
    # Optional: Freeze base layers if you want to fine-tune only top layers
    # for param in policy.model.base_model.parameters():
    #     param.requires_grad = False

# Pass this initialization function to your tune.run config
results = tune.run(
    'PPO',
    config={
        **config1,
        "before_train_fn": load_pretrained_weights
    },
    verbose=0,
    stop=stop,
    checkpoint_at_end=True,
    metric='episode_reward_mean',
    mode="max",
    checkpoint_freq=1
)

If you need to resume training exactly from a previous checkpoint (including optimizer state), stick with RLlib’s native checkpoint restore.


3. How do I load a full PyTorch model into a PPO Agent for inference?

You can initialize a PPO Agent, then overwrite its policy’s model weights with your saved PyTorch state dict:

import torch
from ray.rllib.algorithms.ppo import PPOTrainer

# Initialize the agent with your config
agent = PPOTrainer(config=config1, env=env)
# Get the agent's policy and load the saved model weights
policy = agent.get_policy()
policy.model.load_state_dict(torch.load("./best_pytorch_model.pt"))
policy.model.eval()  # Set to evaluation mode

# Now use the agent for inference
obs = env.reset()
action = agent.compute_action(obs)

For even lighter inference (without initializing the full agent), you can directly load the model:

import torch
from ray.rllib.algorithms.ppo.ppo_torch_policy import PPOTorchPolicy

# Initialize just the policy (no full agent)
policy = PPOTorchPolicy(
    observation_space=env.observation_space,
    action_space=env.action_space,
    config=config1
)
model = policy.model
model.load_state_dict(torch.load("./best_pytorch_model.pt"))
model.eval()

# Run inference
obs = env.reset()
obs_tensor = torch.tensor([obs], dtype=torch.float32)
with torch.no_grad():
    logits, _ = model(obs_tensor, None, None)
    action = torch.argmax(logits, dim=1).item()

Fixing Inference OOM & Slow Checkpoint Restores

Your issues with loading multiple models and slow restores stem from RLlib checkpoints being heavy (they store more than just model weights). Here are actionable fixes:

1. Use Standalone PyTorch Models Instead of Checkpoints

As shown earlier, saving just the PyTorch state dict reduces file size drastically, making loads faster and using less memory than restoring full RLlib checkpoints.

2. Model Quantization

Reduce memory usage and speed up inference with PyTorch quantization:

import torch

# Apply static quantization (requires a calibration step with sample data)
model.qconfig = torch.quantization.get_default_qconfig('fbgemm')
torch.quantization.prepare(model, inplace=True)

# Run calibration with a few batches of observations
# for _ in range(10):
#     obs = env.reset()
#     obs_tensor = torch.tensor([obs], dtype=torch.float32)
#     model(obs_tensor, None, None)

torch.quantization.convert(model, inplace=True)

3. Distributed Model Inference with Ray Actors

Use Ray Actors to host each model in a separate process, avoiding single-process OOM and enabling parallel loading:

import ray
import torch
from ray.rllib.algorithms.ppo.ppo_torch_policy import PPOTorchPolicy
import gym

ray.init(log_to_driver=False, num_cpus=10)  # Adjust based on your resources

@ray.remote(num_cpus=1)
class InferenceActor:
    def __init__(self, config, model_path):
        self.env = gym.make(config["env"])
        self.policy = PPOTorchPolicy(
            observation_space=self.env.observation_space,
            action_space=self.env.action_space,
            config=config
        )
        self.model = self.policy.model
        self.model.load_state_dict(torch.load(model_path))
        self.model.eval()
    
    def get_action(self, obs):
        obs_tensor = torch.tensor([obs], dtype=torch.float32)
        with torch.no_grad():
            logits, _ = self.model(obs_tensor, None, None)
            return torch.argmax(logits, dim=1).item()

# Initialize actors for all your models
model_paths = ["./model1.pt", "./model2.pt", ..., "./model10.pt"]
actors = [InferenceActor.remote(config1, path) for path in model_paths]

# Run inference on any model
obs = env.reset()
action = ray.get(actors[0].get_action.remote(obs))

4. Convert to ONNX for Faster, Lighter Inference

Export your PyTorch model to ONNX format and use ONNX Runtime for low-memory, fast inference:

# Export PyTorch model to ONNX
dummy_obs = torch.tensor([env.reset()], dtype=torch.float32)
torch.onnx.export(
    model,
    dummy_obs,
    "./model.onnx",
    input_names=["obs"],
    output_names=["logits"],
    opset_version=11
)

# Inference with ONNX Runtime
import onnxruntime as ort

ort_session = ort.InferenceSession("./model.onnx")
obs = env.reset()
logits = ort_session.run(["logits"], {"obs": [obs]})[0]
action = logits.argmax(axis=1)[0]

内容的提问来源于stack exchange,提问作者Dr. GUO

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.30 13:18:11