Repository Wiki
bytedance/piano_transcription

Cog 与 Replicate 云端部署

本页介绍 piano_transcription 项目如何通过 Cog 打包为可复现的 GPU 容器镜像,并发布到 Replicate 平台提供在线钢琴音频转写服务。核心由两个文件构成:cog.yaml(构建配置)与 predict.py(预测接口)。

Purpose and Scope(目的与范围)

本页覆盖以下内容:

  • cog.yaml 的完整解析:GPU、系统包、Python 版本、Python 依赖与预测入口的定义
  • predict.py 中 Predictor 类的实现细节:setup() 模型加载、predict() 推理与视频可视化流水线
  • Cog 容器在 Replicate 平台上的请求处理生命周期(冷启动 → 预测 → 输出)
  • 部署相关的依赖约束、失败模式与扩展点

以下主题由兄弟页面覆盖,本页不展开:

  • 模型架构与训练流程(pytorch/ 目录),参见模型训练相关页面
  • 本地推理包 piano_transcription_inference 的内部实现,参见本地推理相关页面
  • 损失函数与评测脚本,参见评测相关页面

Overview(概述)

piano_transcription 是字节跳动提出的高分辨率钢琴转写系统的 PyTorch 实现。为了让不具备 GPU 环境的用户也能使用该模型,仓库提供了 Cog 配置,使其可以在 Replicate 平台上以"上传音频 → 返回钢琴演奏可视化视频"的形式对外服务。README 中给出了对应入口:

text
Demo and Docker image on Replicate https://replicate.com/replicate/piano-transcription

Source: README.md

关键概念:

概念说明
CogReplicate 开源的机器学习容器化工具,通过 cog.yaml + predict.py 定义镜像与推理接口
Replicate模型托管平台,直接读取仓库中的 Cog 配置即可构建并托管模型
PredictorCog 的预测器基类,setup() 在容器启动时执行一次,predict() 每次请求执行一次
piano_transcription_inference封装了预训练模型的官方推理包(版本 0.0.5),部署时通过 pip 安装而非直接使用仓库训练代码
synthviz基于 MIDI 生成钢琴演奏可视化视频的工具,是本部署的最终输出形式

设计意图:仓库没有把训练代码直接塞进部署镜像,而是依赖独立的 piano_transcription_inference PyPI 包。这样部署镜像体积更小、依赖边界清晰,训练与推理解耦——训练代码变动不会影响线上服务稳定性。

Architecture(架构)

Loading diagram...

图中的关键分层:

  1. 平台层:Replicate 负责镜像构建(由 cog.yaml 驱动)、GPU 调度与 HTTP API 暴露。
  2. 预测器层:predict.py:Predictor 是唯一业务入口,Cog 通过装饰器反射出输入参数 schema。
  3. 推理层:piano_transcription_inference 包执行真正的音频转 MIDI;synthviz 把 MIDI 渲染为视频。
  4. 系统依赖层:ffmpeg/timidity/libsndfile1-dev 等系统包支撑音频解码与 MIDI 合成,这是纯 Python 依赖无法覆盖的部分。

文件构成

部署能力由仓库根目录的两个文件完全定义:

文件角色
cog.yaml容器构建配置:GPU、系统包、Python 版本与依赖、入口声明
predict.pyCog 预测器:模型加载、推理、可视化与输出

cog.yaml 的最后一行是入口声明,将预测器类与文件绑定:

yaml
predict: "predict.py:Predictor"

Source: cog.yaml

这行配置告诉 Cog:构建出的容器在接收预测请求时,实例化 predict.py 中的 Predictor 类并调用其 predict 方法。

核心流程:Cog 预测器实现

完整预测器代码

python
1import os 2from pathlib import Path 3 4import cog 5import librosa 6 7# model repo: https://github.com/bytedance/piano_transcription 8# package repo: https://github.com/qiuqiangkong/piano_transcription_inference 9from piano_transcription_inference import PianoTranscription, sample_rate 10from synthviz import create_video 11 12# adapted from example: https://github.com/minzwon/sota-music-tagging-models/blob/master/predict.py 13 14 15class Predictor(cog.Predictor): 16 transcriptor: PianoTranscription 17 18 def setup(self): 19 self.transcriptor = PianoTranscription( 20 device="cuda", checkpoint_path="./model.pth" 21 ) 22 23 @cog.input("audio_input", type=Path, help="Input audio file") 24 def predict(self, audio_input): 25 midi_intermediate_filename = "transcription.mid" 26 video_filename = os.path.join(Path.cwd(), "output.mp4") 27 audio, _ = librosa.core.load(str(audio_input), sr=sample_rate) 28 # Transcribe audio 29 self.transcriptor.transcribe(audio, midi_intermediate_filename) 30 31 # 'Visualization' output option 32 create_video( 33 input_midi=midi_intermediate_filename, video_filename=video_filename 34 ) 35 print( 36 f"Created video of size {os.path.getsize(video_filename)} bytes at path {video_filename}" 37 ) 38 # Return path to video 39 return Path(video_filename)

Source: predict.py

生命周期:setup() 与 predict()

Cog 的预测器遵循两条明确的生命周期钩子,二者的执行次数与时序不同:

方法执行时机执行次数本项目中的职责
setup()容器启动、首次请求到达前一次加载 model.pth 权重到 CUDA 设备,构建 PianoTranscription 实例
predict()每次预测请求每请求一次音频加载 → 转写 → 视频渲染 → 返回输出路径

setup() 的设计意图:

python
1def setup(self): 2 self.transcriptor = PianoTranscription( 3 device="cuda", checkpoint_path="./model.pth" 4 )

Source: predict.py

模型权重加载是最耗时的初始化步骤(涉及磁盘读取与 GPU 显存分配)。将其放在 setup() 而非 predict() 中,保证权重只加载一次、常驻显存,后续请求只需执行推理,避免每次请求重复加载。device="cuda" 硬编码依赖 cog.yaml 中 gpu: true 的构建配置,二者必须保持一致。类属性声明 transcriptor: PianoTranscription 提供了类型标注,便于静态检查。

predict() 的四个阶段:

python
midi_intermediate_filename = "transcription.mid" video_filename = os.path.join(Path.cwd(), "output.mp4") audio, _ = librosa.core.load(str(audio_input), sr=sample_rate)

Source: predict.py

第一阶段是输入解码。librosa.core.load 把任意格式音频(mp3/wav/flac 等,由 ffmpeg/libsndfile 解码)重采样到 sample_rate——该常量从 piano_transcription_inference 包导入,即模型训练时使用的采样率(16 kHz),保证输入分布与训练一致。丢弃的第二个返回值是原始采样率。

python
1self.transcriptor.transcribe(audio, midi_intermediate_filename) 2create_video( 3 input_midi=midi_intermediate_filename, video_filename=video_filename 4)

Source: predict.py

第二、三阶段是转写与可视化:transcribe() 输出中间 MIDI 文件 transcription.mid,随后 synthviz.create_video() 以该 MIDI 为输入渲染演奏视频 output.mp4。中间文件写在容器工作目录中(相对路径 transcription.mid),Cog 容器以临时运行目录为工作目录,属可写空间。

返回值是 Path(video_filename)——Cog 约定 Path 类型返回值会被自动上传并作为可下载产物返回给调用方,这就是用户最终拿到的钢琴演奏视频。

输入参数声明

python
@cog.input("audio_input", type=Path, help="Input audio file") def predict(self, audio_input):

Source: predict.py

@cog.input 装饰器在构建时被反射解析,生成 Replicate API 的输入 schema:

  • name="audio_input":API 参数名,调用方以此字段上传音频
  • type=Path:文件类型输入。请求中的音频会被 Cog 下载到容器内临时路径,以 Path 传给函数
  • help="Input audio file":生成 API 文档时展示给用户的字段说明

因此整条链路的输入是一个音频文件,输出是一个视频文件,没有其他可调参数(如设备选择、阈值等)——这是刻意的极简接口设计,降低线上误用风险。

核心时序图

Loading diagram...

时序要点:

  1. 冷启动成本集中在 setup() 的模型加载上;容器保持热状态时该步骤被跳过。
  2. 每次请求的耗时主体是 GPU 推理(transcribe)与视频渲染(create_video),后者依赖 ffmpeg/timidity 系统包做音频合成。
  3. Cog 自动处理请求中的文件下载(输入)与产物上传(输出),predict() 本身不感知网络。

构建配置详解(cog.yaml)

yaml
1# Configuration for Cog ⚙️ 2# Reference: https://github.com/replicate/cog/blob/main/docs/yaml.md 3 4build: 5 gpu: true 6 7 system_packages: 8 - "libgl1-mesa-glx" 9 - "libglib2.0-0" 10 - "libsndfile1-dev" 11 - "ffmpeg" 12 - "timidity" 13 14 python_version: "3.8" 15 16 python_packages: 17 - "torch==1.8.0" 18 - "torchvision==0.9.0" 19 - "piano_transcription_inference==0.0.5" 20 - "librosa==0.6.0" 21 - "h5py==2.10.0" 22 - "pandas==1.1.2" 23 - "librosa==0.6.0" 24 - "numba==0.48" 25 - "mido==1.2.9" 26 - "mir_eval==0.5" 27 - "matplotlib==3.0.3" 28 - "torchlibrosa==0.0.4" 29 - "sox==1.4.0" 30 - "tqdm==4.62.3" 31 - "pretty_midi==0.2.9" 32 - "synthviz==0.0.2" 33 34 run: 35 - "ffmpeg -version" 36 37predict: "predict.py:Predictor"

Source: cog.yaml

配置项逐项解析

build.gpu: true — 声明镜像需要 GPU。Replicate 据此调度 GPU 实例;predict.py 中的 device="cuda" 依赖此配置成立。

build.system_packages — 五个 apt 系统包,各自不可省略的原因:

包作用被谁使用
ffmpeg音频/视频编解码(mp3 解码、mp4 编码)librosa 音频加载、synthviz 视频渲染
timidityMIDI 合成为音频(波形生成)synthviz 可视化流程
libsndfile1-devlibsndfile 音频读写库librosa/soundfile 读取 wav/flac
libgl1-mesa-glxOpenGL 运行库matplotlib 等绘图依赖的底层库
libglib2.0-0GLib 共享库librosa 依赖链(audioread 等)需要

build.python_version: "3.8" — 锁定 Python 3.8。与 README 声明的开发环境(Python 3.7)略有差异,属于 Cog 基础镜像与依赖版本的折中选择。

build.python_packages — 全部精确钉死版本(==),保证镜像可复现构建。关键条目:

  • piano_transcription_inference==0.0.5:模型推理本体,内部封装了网络结构与权重下载
  • synthviz==0.0.2:视频渲染
  • torch==1.8.0 / torchvision==0.9.0:与推理包兼容的 PyTorch 组合
  • librosa==0.6.0 + numba==0.48:老版本组合,numba 版本必须与 Python/numpy 版本严格匹配,钉版本避免兼容性崩溃
  • mido/pretty_midi:MIDI 文件读写(transcribe 的输出与 create_video 的输入均涉及 MIDI 解析)
  • 其余(h5py/pandas/mir_eval/matplotlib/torchlibrosa/sox/tqdm)为推理包与工具链的传递依赖,显式声明以固化版本
  • 注意 librosa==0.6.0 在列表中出现了两次(L20 与 L23),pip 会忽略重复项,行为无影响,但属于配置冗余

build.run — 构建期 shell 命令。ffmpeg -version 将 ffmpeg 版本号写入构建日志,作为环境验证手段;即便命令失败也不影响依赖安装语义,是构建时的可观测性措施。

predict: "predict.py:Predictor" — 将构建产物与预测入口绑定,Cog 据此定位 Predictor 类。

配置选项参考

cog.yaml 构建选项

选项类型默认说明
build.gpuboolfalse声明镜像需要 GPU;本仓库设为 true,与 device="cuda" 对应
build.system_packageslist[]apt 系统包列表(ffmpeg、timidity、libsndfile1-dev 等)
build.python_versionstring—(本例 "3.8")容器内 Python 解释器版本
build.python_packageslist[]pip 安装的 Python 包列表,全部 == 钉死版本
build.runlist[]构建期执行的 shell 命令(本例 ffmpeg -version)
predictstring—(本例 predict.py:Predictor)预测器入口 文件:类名

Predictor 接口

setup(self)

参数: 无(仅 self)。

职责: 容器启动后、接收首个请求前调用一次。创建 PianoTranscription(device="cuda", checkpoint_path="./model.pth"),将 model.pth 权重加载到 GPU。

副作用: 设置实例属性 self.transcriptor,供后续所有 predict() 调用复用。

predict(self, audio_input)

参数:

  • audio_input (pathlib.Path,必需):请求中上传的音频文件在容器内的本地路径。由 @cog.input("audio_input", type=Path, help="Input audio file") 声明。

返回值: pathlib.Path,指向生成的 output.mp4。Cog 将该文件上传并作为预测结果返回给调用方。

Throws(失败模式):

  • 音频无法解码时 librosa.core.load 抛出异常(依赖 ffmpeg/libsndfile,若容器内对应系统包缺失则一定失败)
  • GPU 显存不足时 transcribe 抛出 CUDA OOM
  • create_video 依赖 timidity 合成,MIDI 渲染失败将中断整个预测

Usage Examples(使用示例)

在本地通过 Cog 运行

仓库配置了 Cog 后,可在本地构建并运行(需安装 cog CLI 与 NVIDIA GPU 驱动):

bash
cog build -t piano-transcription cog predict -i audio_input=@resources/cut_liszt.mp3

以上命令为 Cog CLI 的标准用法:-t 打标签构建镜像,-i 以 文件上传 形式传入 audio_input 参数(输入名来自 @cog.input("audio_input", ...) 声明)。

本地等价推理(不经过 Cog)

README 给出的本地推理流程与 predict() 内部逻辑一致,可用于理解部署链路:

python
1from piano_transcription_inference import PianoTranscription, sample_rate, load_audio 2 3# Load audio 4(audio, _) = load_audio('resources/cut_liszt.mp3', sr=sample_rate, mono=True) 5 6# Transcriptor 7transcriptor = PianoTranscription(device='cuda') # 'cuda' | 'cpu' 8 9# Transcribe and write out to MIDI file 10transcribed_dict = transcriptor.transcribe(audio, 'cut_liszt.mid')

Source: README.md

区别在于:predict.py 使用 librosa.core.load 而非包内提供的 load_audio,并额外增加了 synthviz.create_video 视频化步骤——部署接口面向普通用户,返回可视化视频比返回 MIDI 文件更直观。

失败模式、边界与并发

失败模式

失败场景触发条件位置现象与处理
权重文件缺失./model.pth 不存在setup()容器启动即失败,请求被平台拒绝(冷启动失败)
CUDA OOM音频过长导致推理显存峰值超限transcribe()单次预测失败,容器仍存活
音频解码失败上传文件损坏或格式不支持librosa.core.load异常向上抛出,预测失败
MIDI 渲染失败timidity 缺失或 MIDI 无效create_video视频未生成,预测失败
版本漂移未钉版本的依赖升级构建期本仓库通过全量 == 钉版本规避

边界情况

  • librosa==0.6.0 与 numba==0.48 的版本强耦合:老版本 librosa 依赖 numba JIT 编译部分函数,numba 与 Python/numpy 的 ABI 必须匹配,这是钉版本最重要的动机之一。
  • 采样率一致性:predict() 把音频重采样到从推理包导入的 sample_rate,属于隐式契约——若模型版本变化导致采样率改变,导入值随之变化,无需改动 predict.py。
  • 中间产物路径:transcription.mid 写在容器工作目录,output.mp4 显式拼接到 Path.cwd()。二者均为相对路径,依赖 Cog 容器的可写工作目录语义。
  • 非钢琴音频输入:模型在 MAESTRO 钢琴数据上训练,输入非钢琴音频不会报错,但输出质量不可预期——这是模型边界而非代码边界。

并发行为

Predictor 是被 Cog 容器管理的单实例对象:setup() 只执行一次,self.transcriptor 在所有请求间共享。源码中没有加锁或实例级隔离——GPU 推理由 piano_transcription_inference 内部处理,predict() 中的中间文件名(transcription.mid、output.mp4)是固定字符串,若同一容器内并发执行两次 predict() 会发生文件覆盖竞态。实际部署中 Replicate 通过容器级隔离(每请求一个容器/排队)规避了这一风险,但这是平台的调度保证,而非代码本身的线程安全。

性能与运维注意事项

  • 冷启动:包含镜像拉取 + setup() 模型加载。model.pth 权重读取是主要开销,权重常驻显存后热请求显著更快。
  • 请求耗时主体:GPU 推理与视频渲染各占一部分;create_video 依赖 CPU 上的 timidity 合成,长音频渲染时间会线性增长。
  • 镜像体积:torch==1.8.0 等重型依赖使镜像偏大;librosa==0.6.0 重复声明(L20/L23)虽无害,清理可减少歧义。
  • 依赖约束:piano_transcription_inference==0.0.5 是部署所用的推理包版本,升级时需同步验证 torch 组合兼容性。

扩展点

  • 新增输入参数:在 predict() 上叠加 @cog.input 装饰器(如输出格式开关、设备选择),Cog 会自动把它暴露到 API schema。当前刻意保持单一 audio_input,是最小接口。
  • 返回 MIDI 文件:把 return Path(video_filename) 改为同时返回 Path("transcription.mid")(需配合 Cog 多输出定义),即可在视频之外提供 MIDI 下载。
  • 替换可视化:synthviz 是独立可替换的渲染层;只要消费 transcription.mid 即可换用其他渲染器,predict() 其余逻辑不变。
  • 本地调试:predict.py 头部注释表明其结构改编自 sota-music-tagging-models 的示例,可作为编写其他 Cog 预测器的参考模板。

相关链接