图像标签生成
更新时间:2026-09-14
简介
RAM++ 图像打标算子:用 RAM++(recognize-anything plus,swin_l backbone)给图像打开放词表标签,每张图输出一组去重保序的标签。默认权重 ram_plus_swin_large_14m.pth。
功能描述
- 标签来自 RAM++
generate_tag输出的"|"分隔串,按分隔符拆开、去空白、去重且保持原顺序 - 先把整列图片解码成 RGB,再按
batch_size分微批推理,而不是逐样本推理,吞吐更高 - 权重只从本地
{model_path}/{model_name}加载,缺失时构造期抛FileNotFoundError;算子内不做网络请求与安装 - 为兼容 RAM++ 自带的旧版 BERT 代码,加载时把
transformers.pytorch_utils的apply_chunking_to_forward/prune_linear_layer/find_pruneable_heads_and_indices按需回填到transformers.modeling_utils(RAM++ 照 transformers 4.2x 写,镜像里是 4.57),避免为一个包降级全环境的 transformers - 解码失败的行单独跳过(记日志、该行留空列表),不影响同批其它行;某个微批推理抛异常时该批所有行都留空列表并记 exception
- 输入为 null 时该行返回空列表
- backbone 固定
swin_l,模型输入分辨率由input_size决定(默认 384) - 设备:
num_gpus > 0且torch.cuda.is_available()时用cuda:(rank % 可见卡数),否则 CPU
算子参数
输入
| 输入 | 含义 |
|---|---|
| image | 图像输入,内容类型由 image_src_type 决定(本地/BOS/HTTP 路径、Base64 字符串、二进制数据) |
输出
| 输出 | 含义 |
|---|---|
| tags | list<large_string>:该图的开放词表标签(去重、保序);输入为空或该行失败时为空列表 |
参数
| 参数名称 | 类型 | 默认值 | 描述 |
|---|---|---|---|
| image_src_type | str | "image_url" | 图像输入类型:image_url / image_base64 / image_binary |
| model_path | str | "/opt/aihc/model" | 权重根目录 |
| model_name | str | "ram_plus_swin_large_14m.pth" | 相对 model_path 的 RAM++ 权重文件名 |
| input_size | int | 384 | 模型输入分辨率 |
| batch_size | int | 8 | 模型推理微批大小(<=1 时按 1 张一批) |
| rank | int | 0 | 多卡场景 worker 序号,设备取 cuda:(rank % 可见卡数) |
注意事项
- 需要镜像内已安装
ram(recognize-anything)包,未安装时构造期抛RuntimeError;算子不在运行时装包。 construct_args里的batch_size是模型推理微批,aihc_udf(batch_size=...)是一次 UDF 调用处理多少行,两者是不同的东西:单次调用会先把该次的所有行解码成 PIL 图(内存占用按行数走),再按推理微批进显存。input_size要与权重的训练分辨率一致,默认 384;改了会影响标签质量。
调用示例
Python
1from __future__ import annotations
2
3import os
4
5import daft
6from daft import col
7
8from daft.aihc.common.udf import aihc_udf
9from daft.aihc.functions.image.image_tagging import ImageTagging
10
11if __name__ == "__main__":
12 if os.getenv("DAFT_RUNNER", "native") == "ray":
13 import ray
14 ray.init(dashboard_host="0.0.0.0", ignore_reinit_error=True)
15 daft.set_runner_ray()
16 daft.set_execution_config(actor_udf_ready_timeout=6000, min_cpu_per_task=0)
17
18samples = {"image": ["bos://your-bucket/sample.jpg"]}
19 ds = daft.from_pydict(samples)
20 ds = ds.with_column(
21 "tags",
22 aihc_udf(
23 ImageTagging,
24 construct_args={
25 "image_src_type": "image_url",
26 "model_path": "/path/to/models",
27 "model_name": "ram_plus_swin_large_14m.pth",
28 "input_size": 384,
29 "batch_size": 4,
30 },
31 num_cpus=1,
32 num_gpus=1,
33 concurrency=1,
34 batch_size=8,
35 )(col("image")),
36 )
37 ds.show()
评价此篇文章
