How to Load and Use Pretrained RWKV-CLIP Weights for Production Inference
Load pretrained RWKV-CLIP weights by setting required environment variables, initializing the model skeleton with create_RWKV_Model, and running inference through wrapper functions that handle vision and text embeddings with proper pooling.
The deepglint/rwkv-clip repository implements a multimodal architecture that combines a Vision-RWKV backbone with a Text-RWKV encoder to mirror the CLIP framework. Unlike monolithic training artifacts, this implementation separates model definitions from learned parameters, requiring specific initialization steps to properly load and use pretrained RWKV-CLIP weights for inference in production environments.
Architecture Overview
RWKV-CLIP replaces the traditional Transformer blocks found in standard CLIP with RWKV (Receptance Weighted Key Value) layers for both vision and language processing. The vision encoder processes image patches through spatial mixing operations, while the text encoder handles token sequences with adaptive pooling.
Two critical wrapper functions bridge the gap between raw RWKV outputs and CLIP-compatible embeddings:
WarperCLIP_V_T_RWKV_method– Handles the vision pathway, managing conditional CLS token logic and applying GlobalAveragePooling to produce a single image embedding.WarperCLIP_V_T_RWKV_text_change_head– Processes text hidden states through normalization andnn.AdaptiveAvgPool1d(1)to match the original CLIP text encoder behavior.
Three-Step Weight Loading Process
Loading pretrained weights requires precise orchestration of environment configuration, model instantiation, and state dict cleaning.
1. Configure Environment Variables
The vision RWKV blocks require specific CUDA kernel parameters set before model import. Define these in your shell or Python environment:
export Image_T_max=256
export Image_HEAD_SIE=64
These values must match the checkpoint's training configuration. Image_T_max controls the maximum token length for image patches, while Image_HEAD_SIE defines the head size dimension.
2. Initialize the Model Skeleton
Instantiate the model architecture using create_RWKV_Model from model/utils.py (line 14). This function constructs both encoders based on hyperparameters passed via an args object:
from model import create_RWKV_Model
model = create_RWKV_Model(
args,
model_weight_path="/path/to/rwkv_clip.pth"
)
The function automatically calls load_model_weight (defined at line 95 in model/utils.py), which strips module. and _orig_mod. prefixes introduced by torch.nn.DataParallel or DistributedDataParallel during distributed training.
3. Prepare for Production Inference
Move the model to GPU and enable evaluation mode to disable dropout and activate cuDNN optimizations:
model.cuda()
model.eval()
The model now accepts batched inputs through its forward method, which internally routes images through WarperCLIP_V_T_RWKV_method and text through WarperCLIP_V_T_RWKV_text_change_head (lines 46 and 66 in model/utils.py).
Production Inference Implementation
Minimal PyTorch Script
This self-contained snippet demonstrates the complete inference pipeline:
import os
import torch
from model import create_RWKV_Model, WarperCLIP_V_T_RWKV_method, WarperCLIP_V_T_RWKV_text_change_head
from model.open_clip.tokenizer import tokenize
from model.open_clip.transform import image_transform
from PIL import Image
# Required environment setup
os.environ["Image_T_max"] = "256"
os.environ["Image_HEAD_SIE"] = "64"
# Define configuration stub (match your checkpoint's training args)
class Args:
input_size = 224
image_patch_size = 16
image_embed_dims = 384
image_hidden_rate = 4
image_depth = 12
image_num_heads = 6
image_output_cls_token = False
image_with_cls_token = False
data_type = "utf-8"
ctx_len = 77
vocab_size = 49408
text_initialization = True
head_size = 64
text_num_head = 0
head_size_divisor = 8
n_layer = 12
n_embd = 384
dim_att = 0
dim_ffn = 0
pre_ffn = 0
pos_emb = 0
head_qk = 0
tiny_att_dim = 0
tiny_att_layer = -999
args = Args()
# Build and load model
model = create_RWKV_Model(args, model_weight_path="/path/to/rwkv_clip.pth")
model.eval().cuda()
# Preprocess inputs
transform = image_transform(args.input_size, False)
raw_image = Image.open("example.jpg").convert("RGB")
image_tensor = transform(raw_image).unsqueeze(0).cuda()
texts = ["a photo of a cat."]
token_ids = tokenize(texts).cuda()
# Forward pass returns (image_emb, text_emb, logit_scale)
image_emb, text_emb, logit_scale = model(image_tensor, token_ids)
# Compute cosine similarity
image_norm = torch.nn.functional.normalize(image_emb, dim=-1)
text_norm = torch.nn.functional.normalize(text_emb, dim=-1)
similarity = logit_scale * image_norm @ text_norm.t()
Zero-Shot Classification Pipeline
For batch classification across multiple categories, use the zero_shot_classifier utility from zero_shot.py (line 59):
import json
from zero_shot import zero_shot_classifier
# Load class definitions and templates
with open('utils/label.json') as f:
class_labels = json.load(f)['imagenet']
with open('utils/template.json') as f:
templates = json.load(f)['imagenet']
# Pre-compute classifier weights (run once)
classifier_weights = zero_shot_classifier(model, class_labels, templates, args)
classifier_weights = classifier_weights.cuda()
# Inference per image
image_emb, _, _ = model(image_tensor, token_ids) # token_ids can be dummy here
image_emb = torch.nn.functional.normalize(image_emb, dim=-1)
logits = image_emb @ classifier_weights
predicted_class = class_labels[logits.argmax().item()]
FastAPI Deployment Example
Deploy as a scalable REST endpoint:
from fastapi import FastAPI, File, UploadFile
from PIL import Image
import io, torch, os
app = FastAPI()
# Initialize once at startup
os.environ["Image_T_max"] = "256"
os.environ["Image_HEAD_SIE"] = "64"
model = create_RWKV_Model(args, model_weight_path="/path/to/rwkv_clip.pth")
model.eval().cuda()
transform = image_transform(args.input_size, False)
@app.post("/similarity")
async def similarity(image: UploadFile = File(...), text: str = "a photo of a cat."):
# Image preprocessing
contents = await image.read()
pil_img = Image.open(io.BytesIO(contents)).convert("RGB")
img_tensor = transform(pil_img).unsqueeze(0).cuda()
# Text tokenization
tokens = tokenize([text]).cuda()
# Forward and normalize
img_emb, txt_emb, scale = model(img_tensor, tokens)
img_emb = torch.nn.functional.normalize(img_emb, dim=-1)
txt_emb = torch.nn.functional.normalize(txt_emb, dim=-1)
score = (scale * img_emb @ txt_emb.t()).item()
return {"similarity_score": score}
Production Optimization Guidelines
Optimize your deployment with these specific configurations:
- GPU Placement – Call
model.cuda()once after loading weights. Usemodel.eval()to ensure deterministic batch normalization and disable training-specific layers. - Precision Modes – The repository supports
fp16,bf16, andfp32. Cast tensors using.half()or.bfloat16()for reduced memory bandwidth, or specify--precisionflags if wrapping the training scripts. - Batch Processing – The vision encoder is fully convolutional, allowing variable batch sizes up to GPU memory limits. Process multiple images per call to maximize throughput.
- Determinism – Reuse the
setup_seedhelper from the training utilities if reproducible outputs are required, though this is optional for standard inference. - Checkpoint Compatibility – The
load_model_weightfunction automatically handles checkpoints saved underDataParallelwrappers by removing prefix strings, ensuring compatibility regardless of training distribution strategy.
Summary
- Set environment variables
Image_T_maxandImage_HEAD_SIEbefore importing the model to configure vision block parameters. - Initialize with
create_RWKV_Modelfrommodel/utils.py, passing your configuration args and checkpoint path to handle both skeleton creation and weight loading. - Enable production mode by calling
model.cuda()andmodel.eval()immediately after initialization. - Process inputs through the model's forward method, which automatically routes vision data through
WarperCLIP_V_T_RWKV_methodand text throughWarperCLIP_V_T_RWKV_text_change_head. - Normalize embeddings and apply the returned
logit_scaleto compute CLIP-style similarity scores, or usezero_shot_classifierfor efficient batch classification.
Frequently Asked Questions
Where are the pretrained RWKV-CLIP weights stored?
The deepglint/rwkv-clip repository contains only model architecture definitions. Pretrained learned weights are distributed as separate checkpoint files (e.g., rwkv_clip.pth) that you download independently and load via the model_weight_path argument in create_RWKV_Model.
Why must I set Image_T_max and Image_HEAD_SIE environment variables?
These variables configure the CUDA kernels and head size dimensions for the Vision-RWKV blocks at import time. The vision encoder in model/Image_rwkv.py reads these values during module initialization to allocate appropriate tensor shapes for spatial mixing operations.
How does RWKV-CLIP handle checkpoints saved with DataParallel?
The load_model_weight function in model/utils.py (line 95) automatically cleans state dictionaries by stripping module. and _orig_mod. prefixes. This ensures checkpoints saved during distributed training with torch.nn.DataParallel or DistributedDataParallel load seamlessly into single-GPU or non-distributed multi-GPU inference environments.
Can I run RWKV-CLIP inference in mixed precision?
Yes. The model architecture supports floating-point formats including fp16, bf16, and fp32. You can manually cast the model and inputs using .half() or .bfloat16(), or specify precision modes through the --precision command-line flag when adapting the provided training scripts for inference services.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →