Image-Text-to-Text
Transformers
Safetensors
qwen3_5
vllm
video
multimodal
reinforcement-learning
temporal-grounding
object-tracking
video-segmentation
visual-question-answering
spatial-reasoning
qwen3.5
conversational
Instructions to use OraRL/Video-ORA-4B with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use OraRL/Video-ORA-4B with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("image-text-to-text", model="OraRL/Video-ORA-4B") messages = [ { "role": "user", "content": [ {"type": "image", "url": "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/p-blog/candy.JPG"}, {"type": "text", "text": "What animal is on the candy?"} ] }, ] pipe(text=messages)# Load model directly from transformers import AutoProcessor, AutoModelForMultimodalLM processor = AutoProcessor.from_pretrained("OraRL/Video-ORA-4B") model = AutoModelForMultimodalLM.from_pretrained("OraRL/Video-ORA-4B", device_map="auto") messages = [ { "role": "user", "content": [ {"type": "image", "url": "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/p-blog/candy.JPG"}, {"type": "text", "text": "What animal is on the candy?"} ] }, ] inputs = processor.apply_chat_template( messages, add_generation_prompt=True, tokenize=True, return_dict=True, return_tensors="pt", ).to(model.device) outputs = model.generate(**inputs, max_new_tokens=40) print(processor.decode(outputs[0][inputs["input_ids"].shape[-1]:])) - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use OraRL/Video-ORA-4B with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "OraRL/Video-ORA-4B" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "OraRL/Video-ORA-4B", "messages": [ { "role": "user", "content": [ { "type": "text", "text": "Describe this image in one sentence." }, { "type": "image_url", "image_url": { "url": "https://cdn.britannica.com/61/93061-050-99147DCE/Statue-of-Liberty-Island-New-York-Bay.jpg" } } ] } ] }'Use Docker
docker model run hf.co/OraRL/Video-ORA-4B
- SGLang
How to use OraRL/Video-ORA-4B with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "OraRL/Video-ORA-4B" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "OraRL/Video-ORA-4B", "messages": [ { "role": "user", "content": [ { "type": "text", "text": "Describe this image in one sentence." }, { "type": "image_url", "image_url": { "url": "https://cdn.britannica.com/61/93061-050-99147DCE/Statue-of-Liberty-Island-New-York-Bay.jpg" } } ] } ] }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "OraRL/Video-ORA-4B" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "OraRL/Video-ORA-4B", "messages": [ { "role": "user", "content": [ { "type": "text", "text": "Describe this image in one sentence." }, { "type": "image_url", "image_url": { "url": "https://cdn.britannica.com/61/93061-050-99147DCE/Statue-of-Liberty-Island-New-York-Bay.jpg" } } ] } ] }' - Docker Model Runner
How to use OraRL/Video-ORA-4B with Docker Model Runner:
docker model run hf.co/OraRL/Video-ORA-4B
File size: 15,270 Bytes
0185029 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 | # Copyright 2024 Bytedance Ltd. and/or its affiliates
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""
PPO config
"""
import os
from dataclasses import asdict, dataclass, field, fields, is_dataclass
from typing import Optional, Tuple
from ..utils.multimodal_contract import normalize_video_source_mode
from ..utils.py_functional import get_abs_path
from ..workers.config import WorkerConfig
def recursive_post_init(dataclass_obj):
if hasattr(dataclass_obj, "post_init"):
dataclass_obj.post_init()
for attr in fields(dataclass_obj):
if is_dataclass(getattr(dataclass_obj, attr.name)):
recursive_post_init(getattr(dataclass_obj, attr.name))
@dataclass
class DataConfig:
train_files: str = ""
val_files: str = ""
prompt_key: str = "prompt"
answer_key: str = "answer"
image_key: str = "images"
video_key: str = "videos"
image_dir: Optional[str] = None
video_fps: float = 2.0
video_max_frames: int = 128
max_prompt_length: int = 512
max_response_length: int = 512
rollout_batch_size: int = 512
mini_rollout_batch_size: Optional[int] = None
val_batch_size: int = -1
format_prompt: Optional[str] = None
override_chat_template: Optional[str] = None
enable_thinking: bool = False
"""forwarded to the tokenizer/processor chat template; False preserves no-CoT runs"""
response_prefix: str = ""
"""literal assistant prefix appended after the generation prompt, e.g. '<think>\\n'"""
shuffle: bool = True
seed: int = 1
min_pixels: Optional[int] = 262144
max_pixels: Optional[int] = 4194304
image_min_pixels: Optional[int] = None
image_max_pixels: Optional[int] = None
video_min_pixels: Optional[int] = None
video_max_pixels: Optional[int] = None
video_total_pixels: Optional[int] = None
val_video_fps: Optional[float] = None
val_video_max_frames: Optional[int] = None
val_video_min_pixels: Optional[int] = None
val_video_max_pixels: Optional[int] = None
val_video_total_pixels: Optional[int] = None
filter_overlong_prompts: bool = True
filter_overlong_prompts_workers: int = 16
use_preprocessed_videos: bool = True
"""deprecated compatibility flag; prefer video_source_mode"""
video_source_mode: Optional[str] = None
"""video source policy: prefer_preprocessed, preprocessed_only, realtime_only"""
preprocessed_video_dir: Optional[str] = None
"""directory containing training preprocessed video files (.pt)"""
val_preprocessed_video_dir: Optional[str] = None
"""directory containing validation preprocessed video files (.pt); defaults to preprocessed_video_dir"""
val_video_source_mode: Optional[str] = None
"""validation video source policy; defaults to video_source_mode"""
inline_video_tensors: bool = False
"""if True, decode the video once in the dataset and pass the decoded frames/metadata
inline through multi_modal_data so vLLM rollout and FSDP worker forward passes do
not re-decode the same mp4 (or reload the same .pt). Saves CPU at the cost of
larger pickled batches; recommended for realtime-decode runs."""
group_by_task: bool = False
"""if True, each batch contains samples from a single task type only
(determined by the `group_by_task_key` field in the JSONL). Avoids mixing
modalities (image vs video) within a batch and concentrates gradients."""
group_by_task_key: str = "problem_type"
"""JSONL field name used to identify the task type for task-grouped batching."""
dataloader_num_workers: int = 8
"""number of subprocesses used by the training/validation StatefulDataLoader.
Controls how many videos are decoded in parallel per training step."""
def post_init(self):
self.image_dir = get_abs_path(self.image_dir, prompt="Image directory")
self.format_prompt = get_abs_path(self.format_prompt, prompt="Format prompt file")
self.override_chat_template = get_abs_path(self.override_chat_template, prompt="Chat template file")
self.preprocessed_video_dir = get_abs_path(self.preprocessed_video_dir, prompt="Preprocessed video directory")
self.val_preprocessed_video_dir = get_abs_path(self.val_preprocessed_video_dir, prompt="Validation preprocessed video directory")
if self.image_min_pixels is None:
self.image_min_pixels = self.min_pixels
if self.image_max_pixels is None:
self.image_max_pixels = self.max_pixels
if self.video_min_pixels is None:
self.video_min_pixels = self.min_pixels
if self.video_max_pixels is None:
self.video_max_pixels = self.max_pixels
if self.val_video_fps is None:
self.val_video_fps = self.video_fps
if self.val_video_max_frames is None:
self.val_video_max_frames = self.video_max_frames
if self.val_video_min_pixels is None:
self.val_video_min_pixels = self.video_min_pixels
if self.val_video_max_pixels is None:
self.val_video_max_pixels = self.video_max_pixels
if self.val_video_total_pixels is None:
self.val_video_total_pixels = self.video_total_pixels
if self.val_preprocessed_video_dir is None:
self.val_preprocessed_video_dir = self.preprocessed_video_dir
self.video_source_mode = normalize_video_source_mode(
self.video_source_mode,
use_preprocessed_videos=self.use_preprocessed_videos,
)
if self.val_video_source_mode is None:
self.val_video_source_mode = self.video_source_mode
else:
self.val_video_source_mode = normalize_video_source_mode(
self.val_video_source_mode,
use_preprocessed_videos=self.use_preprocessed_videos,
)
@dataclass
class AlgorithmConfig:
name: str = "grpo"
"""Algorithm selector; the released recipes are ``grpo`` and ``orarl``."""
gamma: float = 1.0
"""discount factor for ppo gae advantage estimator"""
lam: float = 1.0
"""lambda value for ppo gae advantage estimator"""
adv_estimator: str = "grpo"
"""advantage estimator, support `gae`, `grpo`, `reinforce_plus_plus`, `remax`, `rloo`"""
scale_rewards: bool = True
"""Whether GRPO divides each group-centered outcome reward by the group's
reward standard deviation. True preserves the historical normalized GRPO
objective. False uses the TRL raw-centered objective
`A_i = r_i - mean_group(r)`, avoiding amplification of tiny reward gaps."""
disable_kl: bool = False
"""disable reference model"""
use_kl_loss: bool = False
"""use kl loss instead of kl in reward"""
kl_penalty: str = "kl"
"""kl penalty type, support `kl`, `abs`, `mse`, `low_var_kl`, `full`"""
kl_coef: float = 1e-3
"""kl coefficient"""
kl_type: str = "fixed"
"""kl controller type, support `fixed`, `adaptive`"""
kl_horizon: float = 10000.0
"""kl horizon for adaptive kl controller"""
kl_target: float = 0.1
"""target kl for adaptive kl controller"""
# OraRL stages, in the order the paper applies them.
oracle_injection: bool = False
"""Append one annotation-derived oracle response to each rollout group."""
oracle_injection_mode: str = "append"
"""Oracle insertion mode. The released OraRL recipe requires ``append``."""
oracle_builder: Optional[str] = None
"""Dotted ``module:function`` that constructs an oracle response."""
oracle_replace_index: int = -1
"""Compatibility slot for non-released replacement mode."""
oracle_log_exclude: bool = True
"""Exclude oracle rows from on-policy reward/accuracy logging."""
directional_gain: bool = False
"""Enable the oracle-gap directional gain on on-policy utilities."""
directional_gain_gamma: float = 0.25
"""Exponent for the clipped oracle-gap scale."""
directional_gain_positive_only: bool = True
"""Amplify only utilities pointing toward the positive oracle."""
directional_gain_recenter: bool = True
"""Re-center transformed on-policy utilities before selection."""
detached_oracle_advantage: bool = False
"""Overwrite the oracle row with a detached positive anchor."""
detached_oracle_advantage_scale: float = 2.0
"""Base magnitude of the detached oracle anchor."""
detached_oracle_use_directional_gain: bool = False
"""Whether the policy directional gain also scales the oracle anchor."""
detached_oracle_match_best_ratio: float = 1.2
"""Cap the oracle anchor relative to the strongest positive policy row."""
detached_oracle_match_best_min: float = 0.05
"""Lower bound for the adaptive oracle cap."""
detached_oracle_match_best_max: float = 1.0
"""Upper bound for the adaptive oracle cap."""
oracle_reward_gate_beta: float = 2.0
"""Exponent for the normalized oracle reward-gap gate."""
selection_prune_ratio: float = 0.0
"""Fraction of generated policy rows removed before actor backpropagation."""
selection_keep_oracle: bool = True
"""Force-retain the detached oracle row."""
selection_positive_quota: int = 0
"""Number of positive-advantage policy rows retained per group."""
selection_negative_quota: int = 0
"""Number of negative-advantage policy rows retained per group."""
selection_strict_sign_balance: bool = True
"""Select from true positive/negative buckets with deterministic fallback."""
post_selection_recenter: bool = False
"""Restore zero mean over rows that actually receive gradients."""
post_selection_rms_match: bool = False
"""Only downscale active RMS toward its pre-selection policy reference."""
post_selection_rms_min_scale: float = 0.25
"""Lower bound for post-selection RMS downscaling."""
@dataclass
class TrainerConfig:
total_epochs: int = 15
"""total epochs for training"""
max_steps: Optional[int] = None
"""max steps for training, if specified, total_epochs is ignored"""
project_name: str = "orarl"
"""project name for logger"""
experiment_name: str = "demo"
"""experiment name for logger"""
logger: Tuple[str] = ("console", "wandb")
"""logger type, support `console`, `mlflow`, `swanlab`, `tensorboard`, `wandb`"""
nnodes: int = 1
"""number of nodes for training"""
n_gpus_per_node: int = 8
"""number of gpus per node for training"""
critic_warmup: int = 0
"""critic warmup steps"""
val_freq: int = -1
"""validation frequency, -1 means no validation"""
val_before_train: bool = True
"""validate before training"""
val_only: bool = False
"""validate only, skip training"""
val_generations_to_log: int = 0
"""number of generations to log for validation"""
save_freq: int = -1
"""save frequency, -1 means no saving"""
save_limit: int = -1
"""max number of checkpoints to save, -1 means no limit"""
save_model_only: bool = False
"""save model only, no optimizer state dict"""
keep_optim_only_latest: bool = False
"""when True, only the latest checkpoint keeps optimizer / extra_state / dataloader;
older checkpoints retained within `save_limit` are thinned down to model weights only.
Mutually exclusive with `save_model_only` (when `save_model_only=True` optimizer is
never saved at all, so this flag is a no-op). Useful for saving disk while still
allowing resume from the latest step."""
save_checkpoint_path: Optional[str] = None
"""save checkpoint path, if not specified, use `checkpoints/project_name/experiment_name`"""
load_checkpoint_path: Optional[str] = None
"""load checkpoint path"""
ray_timeline: Optional[str] = None
"""file to save ray timeline"""
find_last_checkpoint: bool = True
"""automatically find the last checkpoint in the save checkpoint path to resume training"""
keep_best_train_ckpt: bool = False
"""When True, additionally save the model-only checkpoint at the step with the
highest *smoothed* training reward into a separate `best_train/` subdirectory.
This sidesteps the rolling `save_limit` window so a peak you discovered mid-run
is not evicted by later (potentially overtrained) saves. No optimizer/dataloader
state is kept, only model weights — purely for downstream evaluation."""
best_train_metric_key: str = "reward/overall"
"""Which key in the per-step `metrics` dict to track. Common choices:
`reward/overall` (default; combines IoU + format), `reward/iou`,
`critic/rewards/mean`."""
best_train_smooth_window: int = 5
"""Number of recent steps to average for the best-train signal. Reduces single-
step noise so we lock onto a true plateau rather than a lucky spike. Set to 1
to disable smoothing."""
best_train_min_step: int = 10
"""Don't track best-train until this many steps have elapsed (skip warmup
where reward jumps quickly and the highest single value is nearly meaningless
for picking a good policy)."""
def post_init(self):
if self.save_checkpoint_path is None:
self.save_checkpoint_path = os.path.join("checkpoints", self.project_name, self.experiment_name)
self.save_checkpoint_path = os.path.abspath(self.save_checkpoint_path) # may be not exist
self.load_checkpoint_path = get_abs_path(self.load_checkpoint_path, prompt="Model checkpoint")
@dataclass
class PPOConfig:
data: DataConfig = field(default_factory=DataConfig)
worker: WorkerConfig = field(default_factory=WorkerConfig)
algorithm: AlgorithmConfig = field(default_factory=AlgorithmConfig)
trainer: TrainerConfig = field(default_factory=TrainerConfig)
def post_init(self):
self.worker.rollout.prompt_length = self.data.max_prompt_length
self.worker.rollout.response_length = self.data.max_response_length
self.worker.rollout.trust_remote_code = self.worker.actor.model.trust_remote_code
self.worker.actor.disable_kl = self.algorithm.disable_kl
self.worker.actor.use_kl_loss = self.algorithm.use_kl_loss
self.worker.actor.kl_penalty = self.algorithm.kl_penalty
self.worker.actor.kl_coef = self.algorithm.kl_coef
# The actor sizes its per-rank mini-batch from the post-selection row
# count k = floor(n*(1-P)) per group rather than the full n.
self.worker.actor.selection_prune_ratio = (
float(self.algorithm.selection_prune_ratio)
if str(self.algorithm.name).strip().lower() == "orarl"
else 0.0
)
def deep_post_init(self):
recursive_post_init(self)
def to_dict(self):
from .orarl_config import public_config_dict
return public_config_dict(self)
|