资讯动态

Cog 训练接口(Training Interface)完整指南:为 Cog 模型定义微调入口

发布时间:2026/9/16 16:10:42 来源:尧图企业网站定制
Cog 训练接口Training Interface完整指南为 Cog 模型定义微调入口【免费下载链接】cogContainers for machine learning项目地址: https://gitcode.com/GitHub_Trending/co/cogCog 的训练 API 允许你为一个已有的 Cog 模型定义微调fine-tuning接口让模型使用者自带训练数据产出派生的微调模型。本文将基于 Cog 官方文档 docs/training.md 完整讲解训练接口的声明方式、输入约束、输出对象、测试方法与底层实现并结合仓库源码CLI、Python SDK、集成测试与示例深入剖析其工作原理。[!WARNING]cog train命令行已被标记为弃用deprecated将在 Cog 的下一个版本中移除。下文介绍的训练 API 仍然可以通过 HTTP API 的/trainings端点使用但 CLI 命令不再推荐用于新项目。从源码看train.go 中该命令的 Cobra 定义已经带有Deprecated标记集成测试 train_deprecated.txtar 也专门验证了Command train is deprecated 与 will be removed in a future version of Cog 两条警告输出。训练接口的工作原理如果你使用过 Cog应该已经熟悉 Runner 类用于定义模型推理接口。Cog 的训练 API 工作方式与之类似定义一个 Python 函数来描述训练过程的输入与输出。输入通常是训练数据、epoch 数、batch size、随机种子等超参数输出通常是一个包含微调后权重的文件。两者的差异点在于推理接口通过run:字段声明训练接口通过train:字段声明都在cog.yaml中推理接口执行的是模型前向推理训练接口执行的是训练循环推理接口的核心方法是run()训练接口的核心方法是train()。最小可运行示例cog.yamlbuild: python_version: 3.13 train: train.py:traintrain.pyfrom cog import File import io def train(param: str) - File: return io.StringIO(hello param)然后这样运行$ cog train -i paramtrain ... $ cat weights hello train其中-i paramtrain将train字符串作为名为param的输入传给train()函数weights是训练输出的默认落盘路径由cog train的-o/--output参数控制默认值为weights。使用类class定义训练接口如果你要运行多次训练并希望省去重复的初始化开销可以使用类来定义训练接口。其工作机制与 Runner 类 完全相同唯一的区别是核心方法名为train。典型场景是在setup()中加载一个体积庞大的基础模型例如 SDXL、Llama 2 的基座然后在每次train()调用时在其之上做增量微调。cog.yamlbuild: python_version: 3.13 train: train.py:Trainertrain.pyfrom cog import File import io class Trainer: def setup(self) - None: self.base_model ... # Load a big base model def train(self, param: str) - File: return self.base_model.train(param) # Train on top of a base model注意类形式的训练接口中setup()是可选方法用于一次性加载基础模型等昂贵操作train()是必需方法。如果使用纯函数形式如上面的train(param: str)则没有setup()阶段每次训练都是独立的。底层实现train 模式如何被调度从源码看训练流程的实际执行链如下cog train命令入口位于 pkg/cli/train.go其Use: train [image]Args: cobra.MaximumNArgs(1)——可以不带参数从当前目录构建并训练也可以传入一个由 Cog 构建好的镜像名直接训练若未传镜像CLI 会先通过model.NewSource(configFilename)读取cog.yaml并调用generateLocalOpenAPISchema(src)静态生成 OpenAPI Schema用于把-i输入的字符串解析成带类型的 Python 参数然后走resolver.Build构建镜像容器启动参数为python -m cog.server.http --x-mode train对应 python/cog/server/http.py 中的--x-mode参数其取值来自 python/cog/mode.py 定义的Mode枚举PREDICT/TRAIN在 train 模式下服务端从环境变量COG_TRAIN_TYPE_STUB读取训练接口的引用该变量由cog build在构建时写入 Dockerfile再交给 Rust 侧的 coglet 服务端coglet.server.serve(..., is_trainTrue)执行训练完成后结果由runPrediction(*predictor, preparedInputs, trainOutPath, true, false)落盘到-o指定的输出路径默认weights。集成测试 train_basic.txtar 验证了完整链路cog train -i n42运行后weights.bin被创建且字节数恰好为 42wc -c weights.bin输出42证明输入被正确解析、训练函数被正确调用、输出权重文件被正确写回。使用Input(**kwargs)定义训练输入使用 Cog 的Input()函数为train()函数的每个参数声明输入元数据描述、默认值、校验约束from cog import Input, Path def train( train_data: Path Input(descriptionHTTPS URL of a file containing training data), learning_rate: float Input(descriptionlearning rate, for learning!, default1e-4, ge0), seed: int Input(descriptionrandom seed to use for training, defaultNone) ) - str: return hello, weightsInput()函数支持的参数与 docs/python.md 中run()的Input()完全一致关键字参数适用类型含义description任意给模型使用者的输入说明default任意输入默认值。未传该参数时输入为必填显式设为None时输入为可选geint/float数值必须大于等于该值leint/float数值必须小于等于该值min_lengthstr字符串最小长度max_lengthstr字符串最大长度regexstr字符串必须匹配该正则表达式choicesstr/int输入的候选值列表train()函数的每个参数都必须带类型注解如str、int、float、bool等。完整支持的类型列表参见 输入与输出类型含Path、Secret、Optional、Union、list、dict等包装类型。Input()不是必须的使用Input()函数能为模型使用者提供更好的文档与校验约束但并非强制要求。你也可以直接用纯 Python 的方式指定默认值或完全省略默认值此时为必填参数def train(self, training_data: str foo bar, # this is valid iterations: int # also valid ) - str: # ...注意training_data: str foo bar中的默认值会在 CLI 与 OpenAPI Schema 中体现等价于Input(defaultfoo bar)而iterations: int没有任何默认值等价于必填的Input()。训练输出Training Output训练输出通常是二进制的权重文件。若需要返回自定义对象或包含多个值的复杂对象可以定义一个TrainingOutput对象继承cog.BaseModel用多个字段承载返回值并用 Python 的-返回类型注解将其声明为train()函数的返回类型from cog import BaseModel, Input, Path class TrainingOutput(BaseModel): weights: Path def train( train_data: Path Input(descriptionHTTPS URL of a file containing training data), learning_rate: float Input(descriptionlearning rate, for learning!, default1e-4, ge0), seed: int Input(descriptionrandom seed to use for training, default42) ) - TrainingOutput: weights_file generate_weights(...) return TrainingOutput(weightsPath(weights_file))要点TrainingOutput继承自 python/cog 中的cog.BaseModel其字段类型支持str、int、float、bool、cog.Path、Optional[T]、list[T]等与 Runner 输出对象 一致weights: Path指向训练生成的权重文件路径。返回后该文件会被 Cog 收集、上传/持久化最终交付给训练调用方输出对象可以包含任意数量的字段例如同时返回weights、metrics、checkpoint等只要字段类型受支持即可。仓库中的示例 examples/hello-train/train.py 展示了一个真实可运行的训练接口train(prefix: str) - TrainingOutput把输入字符串写入output.txt后返回TrainingOutput(weightsPath(output.txt))。其配套的 cog.yaml 同时声明了run: run.py:Runner与train: train.py:train即同一项目可同时具备推理与训练两种接口推理接口 run.py 还演示了通过COG_WEIGHTS在setup()中加载权重的用法。测试训练代码路径如果你正在开发类似 Llama 或 SDXL 这样的 Cog 模型可以在推送之前测试微调代码路径是否正常工作。方法是在运行cog run时通过-e指定COG_WEIGHTS环境变量cog run -e COG_WEIGHTShttps://replicate.delivery/pbxt/xyz/weights.tar -i prompta photo of TOK原理说明有源码依据COG_WEIGHTS环境变量会被 python/cog/predictor.py 中的extract_setup_weights()读取作为weights参数传入setup(weights...)因此只要你的 Runner/Trainer 的setup()方法签名包含weights: Optional[Path] None参数如示例 examples/hello-train/run.py 所示就可以在本地开发时用真实权重 URL 走一遍下载权重 → setup 加载 → 推理/训练的完整路径而无需先推送模型到远端cog run的-e参数用于注入环境变量形式为namevalue与cog train的-e/--env参数行为一致见 train.go。此外训练模式下同样支持cog train -e NAMEvalue -i keyvalue这种组合传参方式-e注入环境变量例如COG_WEIGHTS、CUDA_VISIBLE_DEVICES等-i传递训练输入。若-i的值以开头则从磁盘文件读取内容例如-i train_datadata.json对应 train.go 中if value is prefixed with , then it is read from a file on disk的说明。训练接口与推理接口的对比维度推理接口Runner训练接口Training声明位置cog.yaml的run:字段cog.yaml的train:字段核心方法run(**kwargs)train(**kwargs)可选初始化setup()setup()类形式输入定义Input()函数Input()函数完全一致输出类型任意受支持类型 /BaseModel输出对象通常为权重文件可用BaseModel输出对象运行命令cog run/cog predictcog train已弃用推荐 HTTP/trainings底层模式--x-mode predict--x-mode train见 http.py深入阅读训练接口官方参考本文的直接依据Run 接口参考Runner/Input/输出类型训练接口复用的类型系统与Input()语义hello-train 示例一个同时具备推理与训练接口的最小可运行项目train_basic.txtar 集成测试验证cog train -i输入解析与权重输出落盘的端到端测试train_deprecated.txtar 集成测试验证 CLI 弃用警告训练 CLI 实现cog train命令的完整参数与执行逻辑HTTP 服务入口--x-mode train如何切换到训练模式并读取COG_TRAIN_TYPE_STUB。如果你正在构建一个新的、可被他人微调的模型建议在cog.yaml中同时声明run:与train:用TrainingOutput(BaseModel)返回权重文件并在setup()中预留weights参数以便通过COG_WEIGHTS做本地测试。对于新项目请优先通过 HTTP API 的/trainings端点使用训练能力避免依赖将被移除的cog train命令。【免费下载链接】cogContainers for machine learning项目地址: https://gitcode.com/GitHub_Trending/co/cog创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价