Files
mineru-rocm/docker/scripts/cache_warmer.py
T
2026-06-03 15:39:09 +08:00

97 lines
3.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#!/usr/bin/env python3
"""
ROCm MIOpen Cache Warmer for MinerU
在 AMD RDNA 架构上,MIOpen 遇到新尺寸的卷积运算时需要搜索最优 kernel(冷启动)。
预热脚本提前跑一遍常用尺寸,将 kernel 缓存到 ~/.cache/miopen/,避免运行时等待。
缓存持久化到磁盘,重启不丢失;只有升级 ROCm 后才需要重新跑。
用法:
python cache_warmer.py --device cuda --max_side 960 --step 32
"""
import argparse
import torch
import torch.nn as nn
import torch.nn.functional as F
from tqdm import tqdm
def get_args():
p = argparse.ArgumentParser(description="ROCm MIOpen Cache Warmer")
p.add_argument("--device", type=str, default="cuda")
p.add_argument("--max_side", type=int, default=960)
p.add_argument("--step", type=int, default=32)
return p.parse_args()
class MockOCRModel(nn.Module):
"""模拟 MinerU OCR 模型的卷积结构,覆盖 MIOpen 常用 kernel 尺寸。"""
def __init__(self, in_channels: int = 3):
super().__init__()
self.stem = nn.Conv2d(in_channels, 16, 3, stride=2, padding=1)
self.dw_3x3 = nn.Conv2d(16, 16, 3, stride=1, padding=1, groups=16)
self.pw_1 = nn.Conv2d(16, 64, 1)
self.dw_5x5 = nn.Conv2d(64, 64, 5, stride=2, padding=2, groups=64)
self.pw_2 = nn.Conv2d(64, 128, 1)
self.dw_3x3_s2 = nn.Conv2d(128, 128, 3, stride=2, padding=1, groups=128)
self.pw_3 = nn.Conv2d(128, 256, 1)
self.out_conv = nn.Conv2d(256, 64, 1)
self.binarize_conv = nn.Conv2d(64, 1, 3, stride=1, padding=1)
self.act = nn.ReLU()
def forward(self, x):
x = self.stem(x)
x = self.act(x)
x = self.dw_3x3(x)
x = self.pw_1(x)
x = self.dw_5x5(x)
x = self.act(x)
x = self.pw_2(x)
x = self.dw_3x3_s2(x)
x = self.pw_3(x)
x = self.out_conv(x)
x = F.interpolate(x, scale_factor=2, mode="bilinear", align_corners=True)
x = self.binarize_conv(x)
return x
def main():
args = get_args()
assert torch.cuda.is_available(), "GPU not available"
device = torch.device(args.device)
print("=" * 50)
print("ROCm MIOpen Cache Warmer")
print(f" GPU : {torch.cuda.get_device_name(0)}")
print(f" ROCm : {torch.version.hip}")
print(f" Cache : ~/.cache/miopen/")
print("=" * 50)
model = MockOCRModel().to(device).eval()
sizes = list(range(64, args.max_side + 1, args.step))
combos = [(h, w) for h in sizes for w in sizes]
print(f"Warming {len(combos)} shapes ({len(sizes)}x{len(sizes)} grid, "
f"step={args.step})...")
ok = 0
with torch.no_grad():
for h, w in tqdm(combos, desc="Warming"):
try:
model(torch.zeros((1, 3, h, w), device=device, dtype=torch.float32))
ok += 1
except RuntimeError as e:
if "out of memory" in str(e):
torch.cuda.empty_cache()
# 其他错误跳过,不影响后续
print(f"\nDone! {ok}/{len(combos)} shapes cached (~3–4 min)")
print("Kernels saved to ~/.cache/miopen/")
if __name__ == "__main__":
main()