如何通过代码在Document AI上训练自定义分类器并实现客户自助训练
代码调用Document AI API训练文档分类器解决方案
问题背景
我想通过代码调用Document AI API训练文档分类器,但在官方文档和代码示例里找不到相关信息。已经定义了Invoice OCR处理器,但不知道怎么指定训练集和测试集。需求是让应用里的客户能像在Google平台上一样自行训练处理器,目前的思路是先提交文档处理,从GCS下载生成的JSON文件,修改字段数值和坐标后再用于训练,附上了大致代码:
putenv('GOOGLE_APPLICATION_CREDENTIALS='.$this->parameterBag->get('gmail_private_key')); $client = new DocumentProcessorServiceClient(); $name = $client::processorVersionName(self::PROJECT_ID, self::LOCATION, self::PROCESSOR_ID, self::PROCESSOR_VERSION); $storageClient = new StorageClient(); $outputBlobs = $storageClient->bucket(self::BUCKET_NAME)->objects(['prefix' => $prefix]); $document = new Document(); /** @var StorageObject $blob */ foreach ($outputBlobs as $blob) { // Document AI 应仅输出JSON文件到GCS if ($blob->info()['contentType'] !== "application/json") { continue; } $jsonText = $blob->downloadAsStream(); $fields = json_decode($blob->downloadAsString(), true, 512, JSON_THROW_ON_ERROR); $document->mergeFromJsonString($jsonText); /** @var Document\Entity $entity */ foreach ($document->getEntities() as $entity) { $entity->setType('something'); $entity->setConfidence(0.9); $entity->setPageAnchor('something'); $entity->setNormalizedValue('something'); } } $getProcessorVersionRequest = new GetProcessorVersionRequest(); $getProcessorVersionRequest->setName($name); $processorVersion = $client->getProcessorVersion($getProcessorVersionRequest); $trainProcessorVersionRequest = new TrainProcessorVersionRequest(); $trainProcessorVersionRequest->setProcessorVersion($processorVersion); $gcsDocument = new GcsDocument(); $gcsDocument->setGcsUri(self::GCS_URI); $gcsDocument->setMimeType('application/json'); $gcsDocuments = new GcsDocuments(); $gcsDocuments->setDocuments([$gcsDocument]); $batchDocumentsInputConfig = new BatchDocumentsInputConfig(); $batchDocumentsInputConfig->setGcsDocuments($gcsDocuments); $inputData = new InputData(); $inputData->setTrainingDocuments($batchDocumentsInputConfig); $trainProcessorVersionRequest->setInputData($inputData); $client->trainProcessorVersion($trainProcessorVersionRequest);
解决方案与代码修正
核心逻辑调整说明
你的代码目前是针对实体提取模型的修改逻辑,但文档分类器需要的是文档级的类别标注,不是实体级修改。另外,训练分类器必须用专门的文档分类器处理器,不能用Invoice OCR处理器,下面是修正后的完整流程:
- 先创建文档分类器处理器:通过控制台或API创建一个文档分类器类型的处理器,替换原有Invoice OCR处理器ID。
- 准备标注数据:分类器需要的标注是每个文档对应的类别,推荐用JSONL格式(每行一个JSON),上传到GCS,格式示例:
{"document": {"gcs_uri": "gs://your-bucket/invoice-001.pdf"}, "labels": ["invoice"]} {"document": {"gcs_uri": "gs://your-bucket/receipt-001.pdf"}, "labels": ["receipt"]} - 修正训练请求参数:训练是基于处理器创建新的版本,不需要获取现有版本,同时要指定训练集和测试集。
修正后的代码示例
putenv('GOOGLE_APPLICATION_CREDENTIALS='.$this->parameterBag->get('gmail_private_key')); $client = new DocumentProcessorServiceClient(); // 替换为你的文档分类器处理器ID $processorName = $client::processorName(self::PROJECT_ID, self::LOCATION, self::CLASSIFIER_PROCESSOR_ID); // 配置训练集GCS路径(JSONL格式) $trainingGcsDoc = new GcsDocument(); $trainingGcsDoc->setGcsUri('gs://your-bucket/training-data.jsonl'); $trainingGcsDoc->setMimeType('application/jsonl'); $trainingGcsDocs = new GcsDocuments(); $trainingGcsDocs->setDocuments([$trainingGcsDoc]); $trainingInputConfig = new BatchDocumentsInputConfig(); $trainingInputConfig->setGcsDocuments($trainingGcsDocs); // 配置测试集GCS路径(可选,推荐添加) $testGcsDoc = new GcsDocument(); $testGcsDoc->setGcsUri('gs://your-bucket/test-data.jsonl'); $testGcsDoc->setMimeType('application/jsonl'); $testGcsDocs = new GcsDocuments(); $testGcsDocs->setDocuments([$testGcsDoc]); $testInputConfig = new BatchDocumentsInputConfig(); $testInputConfig->setGcsDocuments($testGcsDocs); // 组装输入数据 $inputData = new InputData(); $inputData->setTrainingDocuments($trainingInputConfig); $inputData->setTestDocuments($testInputConfig); // 配置分类器训练参数(比如类别映射) $classifierParams = new ClassifierTrainingParams(); $classifierParams->setLabelMap(['invoice' => '发票', 'receipt' => '收据']); // 构建训练请求 $trainRequest = new TrainProcessorVersionRequest(); $trainRequest->setParent($processorName); $trainRequest->setInputData($inputData); $trainRequest->setClassifierTrainingParams($classifierParams); // 提交异步训练请求 $operation = $client->trainProcessorVersion($trainRequest); // 等待训练完成 $operation->pollUntilComplete(); // 获取训练好的处理器版本 $trainedVersion = $operation->getResult(); echo "训练完成,新处理器版本名称:" . $trainedVersion->getName();
关键注意事项
- 处理器类型必须匹配:只能用文档分类器处理器,OCR/实体提取处理器无法训练分类模型。
- 标注格式要正确:分类器不识别实体标注,必须提供文档级的类别标签,JSONL是最简便的格式。
- 权限要到位:服务账号需要有GCS的读写权限,以及Document AI的训练权限(roles/documentai.editor)。
- 训练是异步操作:可以通过
pollUntilComplete()等待完成,或者保存操作ID后续查询进度。
内容的提问来源于stack exchange,提问作者Kronchik X
相关产品推荐
相关产品推荐

