房凡鸣头像
关注

使用 Gradio Flagging 标记机制收集模型困难样本:Interface 与 Blocks 全指南

使用 Gradio Flagging 标记机制收集模型困难样本:Interface 与 Blocks 全指南

【免费下载链接】gradio Build and share delightful machine learning apps, all in Python. 🌟 Star to support our work! 【免费下载链接】gradio 项目地址: https://gitcode.com/GitHub_Trending/gr/gradio

导读:在部署机器学习 Demo 时,最值得沉淀的数据往往来自“模型表现不如预期”的样本。Gradio 在每个界面中内置了 Flag(标记) 机制,让试用者一键把输入输出回传、落盘为结构化数据。本文基于官方中文指南《使用标记》(guides/cn/07_other-tutorials/using-flagging.md),结合当前仓库源码,系统讲解 Interface 下控制标记的四个参数、标记数据的落盘格式、自定义 FlaggingCallback 回调,以及在 Blocks 中手动接线标记功能的完整做法。

为什么需要标记功能

当我们把一个机器学习模型做成可交互的演示给用户试用时,通常希望收集用户的真实使用数据,尤其是那些让模型“翻车”的困难样本。捕获这些数据点意义重大:它们是改进模型、提升可靠性与稳健性的直接素材。

Gradio 通过在每个 Interface 输出组件下方内置一个 Flag(标记) 按钮来简化这一流程——试用者或测试人员看到有趣的输入输出组合时,点击按钮即可将数据发送回运行演示的机器。样本默认被写入 CSV 日志文件;若演示涉及图像、音频、视频等文件类型数据,这些文件会被单独保存在并行目录中,而文件路径写入 CSV。

在 Blocks 层面,标记的底层调度由 gradio/flagging.py 中的 FlagMethod 辅助类完成:它包装了回调的 .flag() 调用,并把 gr.Request.username 一并传入;出错时打印 Error while flagging 并返回错误反馈,成功后短暂等待 0.8 秒让用户看到按钮变化,再将其重置为可点击状态。

Interface 中控制标记的四个参数

gradio.Interface 中,标记功能由一组构造参数控制。需要说明的是:中文指南撰写时使用的参数名为 allow_flagging,而在当前仓库版本中,该参数已正式更名为 flagging_mode(见 gradio/interface.py),语义与取值保持一致。下文按当前仓库的实际参数名展开。

1. flagging_mode:标记模式(核心开关)

取值可以是 "manual"(默认)、"auto""never"

取值用户是否看到标记按钮行为
manual仅当用户点击按钮时才标记样本
auto每一次提交都会自动标记(输入与输出一并记录)
never不显示按钮,也不标记任何样本

在源码层面,其解析逻辑位于 gradio/interface.py:优先读取显式传入的参数;若为 None 则回退到环境变量 GRADIO_FLAGGING_MODE;再没有则默认 "manual"。传入非法值会直接抛出 ValueError。当模式为 auto 时,界面启动后每次推理完成都会触发一次标记;当模式为 manual 时,标记按钮的点击事件才会被挂载。

2. flagging_options:让用户为标记提供原因

参数取值为 None(默认)或字符串列表。

  • 若为 None,界面只显示一个 Flag 按钮,点击即标记,无附加选择。
  • 若传入字符串列表,例如 ["wrong sign", "off by one", "other"],界面会渲染出多个按钮,文案为 Flag as wrong signFlag as off by oneFlag as other;用户必须选择其中一项,所选原因会与输入输出一起作为附加列记录。
  • 该参数仅在 flagging_mode="manual" 时生效。

从源码看,flagging_optionsgradio/interface.py 中会被规范化为 (label, value) 二元组列表:纯字符串列表会被转换为 ("Flag as X", "X");也支持直接传入 (label, value) 元组列表,以便按钮文案与 CSV 记录值解耦;传错类型会抛出 ValueError。按钮最终由 _create_flagging_buttons 等内部逻辑逐一渲染。

3. flagging_dir:标记数据的存放目录

接受一个字符串,表示标记数据存储的目录名,目录不存在时会自动创建。当前 Interface 的默认值为 ".gradio/flagged"(见 gradio/interface.py)。注意:在 Hugging Face Spaces 这类临时托管环境中该目录是不可持久化的,因此通常配合自定义回调使用。

4. flagging_callback:标记时的自定义代码入口

接受 FlaggingCallback 子类实例,用于在用户点击标记按钮时执行自定义逻辑。默认不传时为 gr.CSVLogger()gradio/interface.py)。官方文档中的经典用例是传入 HuggingFaceDatasetSaver,把标记数据直接汇入 Hugging Face 数据集以构建“众包”训练集——关于该类的当前状态与自定义写法,详见下文“自定义回调”一节。

标记的数据会发生什么

当用户在计算器示例中点击标记按钮后,启动目录下会生成存放标记数据的子目录。下面是中文指南提供的计算器示例(数字输入 + 四则运算单选 + 数字输出):

import gradio as gr


def calculator(num1, operation, num2):
    if operation == "add":
        return num1 + num2
    elif operation == "subtract":
        return num1 - num2
    elif operation == "multiply":
        return num1 * num2
    elif operation == "divide":
        return num1 / num2


iface = gr.Interface(
    calculator,
    ["number", gr.Radio(["add", "subtract", "multiply", "divide"]), "number"],
    "number",
    flagging_mode="manual"
)

iface.launch()

点击标记后,指南描述的目录结构为:

+-- flagged/
|   +-- logs.csv

logs.csv 的内容形如:

num1,operation,num2,Output,timestamp
5,add,7,12,2022-01-31 11:40:51.093412
6,subtract,1.5,4.5,2022-01-31 03:25:32.023542

如果界面涉及图像、音频等文件数据,例如 image→image 的界面,会额外创建与组件标签对应的子目录来存放文件,路径记录进 CSV:

+-- flagged/
|   +-- logs.csv
|   +-- image/
|   |   +-- 0.png
|   |   +-- 1.png
|   +-- Output/
|   |   +-- 0.png
|   |   +-- 1.png

需要向读者说明的是:上述 flagged/logs.csv 是旧版 Gradio 的落盘形态(指南撰写时的快照)。当前仓库的默认回调 CSVLoggergradio/flagging.py)实际会在 flagging_dir 下生成名为 dataset1.csv 的数据集文件,其行为与旧版有三点明显差异:

  1. 表头按组件标签动态生成:每列依次对应当前参与标记的组件(取组件 label,缺省为 component {idx}),列尾固定追加 timestamp;仅当首次标记时收到 flag_optionusername,才会额外加入 flagusername 列(见 _create_dataset_fileflag 方法的逻辑)。
  2. 并发安全CSVLogger 内部持有 multiprocessing.Lock,写入与行数统计都在锁内完成,flag() 返回截至当前的总标记行数。
  3. 表头变更时自动滚动新文件:当组件标签衍生的表头与最新 dataset*.csv 不一致时,会创建递增编号的新文件(dataset2.csv 等);也支持通过构造参数 dataset_file_name 固定文件名,simplify_file_data 控制是否简化文件路径写入,verbose 控制在建文件时是否打印提示。

测试 test/test_flagging.py 印证了这些行为:flag() 返回带表头的累计行数;文件型组件(如 bus.png)会被保存并以 .png 结尾的相对路径写入 dataset1.csv;纯文本界面标记后目录内只应出现一个 dataset1.csv,不会产生多余文件夹。若需要旧版行为(含 flag/username/timestamp 固定表头),可显式传入 gr.ClassicCSVLogger;还有一个不写表头、仅追加数据的 gr.SimpleCSVLogger 作为最简实现的教学示例(gradio/flagging.py)。

让用户为标记提供原因

回到计算器示例,加入 flagging_options 后,标记得到的 CSV 会新增一列 flag,记录用户选择的原因:

iface = gr.Interface(
    calculator,
    ["number", gr.Radio(["add", "subtract", "multiply", "divide"]), "number"],
    "number",
    flagging_mode="manual",
    flagging_options=["wrong sign", "off by one", "other"]
)

iface.launch()
num1,operation,num2,Output,flag,timestamp
5,add,7,-12,wrong sign,2022-02-04 11:40:51.093412
6,subtract,1.5,3.5,off by one,2022-02-04 11:42:32.062512

自定义回调:从本地 CSV 到任意数据出口

把标记数据写进本地 CSV 并不总是合理。例如在 Hugging Face Spaces 上,开发者通常无法访问托管演示的底层临时机器——这也是为什么在 Spaces 环境中标记默认被关闭(flagging_mode 默认落到 "never")。此时 flagging_callback 参数就是数据出口的扩展点。

指南中的示例是把计算器的标记数据导入一个公共 Hugging Face 数据集以构建众包数据集:先通过 os.getenv("HF_TOKEN") 取得令牌,构造 gr.HuggingFaceDatasetSaver(HF_TOKEN, "crowdsourced-calculator-demo") 并传入 flagging_callback,同时在界面描述中标注数据集的公开地址。

需要基于当前仓库事实向读者澄清两点:

  • 从当前仓库源码结构看,gradio/flagging.py 中已不存在 HuggingFaceDatasetSaver 类,gradio/init.py 也仅导出 CSVLoggerSimpleCSVLoggerFlaggingCallback。指南中该示例描述的是旧版能力;如果你的需求是把标记数据写回云端数据集,请参考你所用 Gradio 版本对应的 API 文档。
  • 不变的核心扩展机制依然完好:任何自定义回调类只要继承 gradio/flagging.py 中定义的抽象基类 FlaggingCallback,并实现两个抽象方法即可被 InterfaceBlocks 使用:
    • setup(components, flagging_dir):在界面启动初期被调用一次,负责为 flag() 做准备(如创建目录、记录组件序列);
    • flag(flag_data, flag_option=None, username=None):每次点击标记按钮时被调用,flag_data 为被标记的输入输出数据,返回当前已标记的总样本数。

这两个方法签名就是“任意数据出口”的契约:你可以把 flag() 里的逻辑替换为写入数据库、推送消息队列、调用标注平台 API 等任何操作。若实现了优质的回调类,指南也鼓励将其贡献回开源仓库。

在 Blocks 中使用标记:两步骤接线

gradio.Blocks 提供了更高的灵活性——你完全可以自己写点击按钮后要执行的 Python 代码并用事件绑定。但如果想复用现成的 FlaggingCallback(如默认的 CSVLogger)来避免重复造轮子,指南给出两个必要步骤:

  1. 在代码中的某个位置(第一次标记发生之前)调用回调的 .setup() 方法;
  2. 点击标记按钮时触发回调的 .flag() 方法,确保正确收集参数、并传入 preprocess=False 禁用常规预处理(因为回调拿到的应当是最接近原始形态的数据)。

仓库示例 demo/blocks_flag/run.py 是一份可直接运行的“图像怀旧滤镜 + 标记”完整代码:

import numpy as np
import gradio as gr

def sepia(input_img, strength):
    sepia_filter = strength * np.array(
        [[0.393, 0.769, 0.189], [0.349, 0.686, 0.168], [0.272, 0.534, 0.131]]
    ) + (1-strength) * np.identity(3)
    sepia_img = input_img.dot(sepia_filter.T)
    sepia_img /= sepia_img.max()
    return sepia_img

callback = gr.CSVLogger()

with gr.Blocks() as demo:
    with gr.Row():
        with gr.Column():
            img_input = gr.Image()
            strength = gr.Slider(0, 1, 0.5)
        img_output = gr.Image()
    with gr.Row():
        btn = gr.Button("Flag")

    # This needs to be called at some point prior to the first call to callback.flag()
    callback.setup([img_input, strength, img_output], "flagged_data_points")

    img_input.change(sepia, [img_input, strength], img_output)
    strength.change(sepia, [img_input, strength], img_output)

    # We can choose which components to flag -- in this case, we'll flag all of them
    btn.click(lambda *args: callback.flag(list(args)), [img_input, strength, img_output], None, preprocess=False)

if __name__ == "__main__":
    demo.launch()

这段代码展示了 Blocks 标记的几个关键点:

  • callback.setup(...) 传入的是你想要记录的全部组件(此处为输入图片、强度滑杆、输出图片三个组件)与目标目录 "flagged_data_points",它决定了 CSV 的列与文件子目录的划分。
  • btn.click(...) 通过一个 lambda 把三个组件的当前值收集为列表交给 callback.flag(list(args))preprocess=False 意味着回调拿到的组件数据未被额外清洗;None 作为输出表示该事件不产生界面返回值。
  • Blocks 里你可以自由选择标记哪些组件,不必像 Interface 那样全量记录——例如只标记输出而忽略滑杆状态也是可行的。

作为横向补充,对话场景也有专门的落盘方案:当前仓库在 gradio/flagging.py 提供了 ChatCSVLogger,将整段对话以 JSON 序列化、连同消息索引与 like/dislike 反馈写入 log.csv,配合 ChatInterface 的点赞/点踩事件使用。

隐私提示

重要提醒:请务必确保用户清楚了解——他们提交的数据何时会被保存,以及你打算如何处理这些数据。这一点在使用 flagging_mode="auto"(即通过演示提交的所有数据都会被自动标记并持久化)时尤其关键。合理的做法是在界面的 descriptionarticle 中明示数据采集策略,避免在用户不知情的情况下收集数据。

小结

标记机制的价值闭环可以概括为三句话:Interface 用四个参数(flagging_modeflagging_optionsflagging_dirflagging_callback)零成本开启数据回流;默认 CSVLogger 负责把样本(含文件型样本的相对路径)可靠地落盘为带表头的 CSV,并保证并发安全与表头变更时的文件滚动;当默认出口不满足需求时,继承 FlaggingCallback 实现 setup()flag(),即可把“困难样本”送到任何你想要的去处。配合 Blocks 的两步接线法,标记能力可以精确地挂在你自定义界面的任意事件上——这正是持续迭代模型所需的最后一公里数据管道。

【免费下载链接】gradio Build and share delightful machine learning apps, all in Python. 🌟 Star to support our work! 【免费下载链接】gradio 项目地址: https://gitcode.com/GitHub_Trending/gr/gradio

转载自 CSDN-专业IT技术社区

原文链接:https://blog.csdn.net/gitblog_00493/article/details/156665641

文章来源转载

评论

赞0

评论列表

微信小程序
QQ小程序

关于作者

点赞数:0
关注数:0
粉丝:0
文章:0
关注标签:0
加入于:--