使用Xenova Transformers生成数据库字段注释出错,求优化方案
数据库字段注释生成优化请求
我基于NestJS框架,使用Xenova Transformers的flan-t5-small模型预测数据库字段注释,部分注释大致正确但多数生成错误。现寻求更优的实现方案,或代码优化指导以提升注释生成的准确率。
以下是完整项目代码、输入输出数据及提示词文件:
column-test-predictor.service.ts
import { Injectable, Logger, OnModuleInit } from '@nestjs/common'; import { pipeline, TextGenerationPipeline, TextGenerationSingle, } from '@xenova/transformers'; import fs from 'fs'; export interface ValuePrediction { values: string[]; confidence: number; method: Method; } enum Method { SIMILARITY = 'similarity', CLASSIFICATION = 'classification', } export interface PredictedColumn { column_name: string; datatype: string; similarity: number; comment: string | null; allowed_values: ValuePrediction; colConfidence: number; } @Injectable() export class ColumnTestPredictorService implements OnModuleInit { private logger: Logger = new Logger(ColumnTestPredictorService.name); private commentGenerator: TextGenerationPipeline; constructor() {} onModuleInit() { this.init(); } async init() { this.commentGenerator = await pipeline( 'text2text-generation', 'Xenova/flan-t5-small', ); this.logger.debug('INIT'); } private async generateComment( columnName: string, table_name: string, ): Promise<string | null> { const basicPrompt = fs .readFileSync( process.cwd() + '/src/models/comment-predictor/prompt.txt', 'utf8', ) .toString(); const fewShotPrompt = ` Guidelines: - Be specific about the content - Mention relationships if applicable - Include measurement units if relevant - Use consistent terminology - Keep it under 15 words unless necessary Examples of inputs: ${basicPrompt} Inputs: Table name: ${table_name} Column name: ${columnName} Output:`; //this.logger.log('fewShotPrompt', fewShotPrompt); const result = await this.commentGenerator(fewShotPrompt, { max_new_tokens: 10, temperature: 0.3, //top_k: 40, //top_p: 0.9, //do_sample: true, }); const [text] = result as TextGenerationSingle[]; const lines = text.generated_text as string; this.logger.debug('text', JSON.stringify(text.generated_text, null, 4)); const rawComment = lines?.trim() ?? null; // Additional cleanup for quality control if (rawComment) { // Ensure the comment is properly formatted return rawComment .replace(/\.$/, '') // Remove trailing period if present .replace(/\s+/g, ' ') // Collapse multiple spaces .trim(); } return null; } async predictColumns( prompt: string, table_name: string, columns: Partial<PredictedColumn>[], ): Promise<Partial<PredictedColumn>[]> { const finalResults: Partial<PredictedColumn>[] = []; let i = 0; for (const result of columns) { finalResults[i] = { ...columns[i], comment: (await this.generateComment( columns[i].column_name as string, table_name, )) as string, }; i++; } return finalResults; } }
database-migrations.controller.ts
import { Body, Controller, Get, Logger, Param, Post, Query } from '@nestjs/common'; import fs, { promises as fsp } from 'fs'; import { HttpException, HttpStatus } from '@nestjs/common'; import { lastValueFrom } from 'rxjs'; import { DatabaseMigrationsService } from './database-migrations.service'; @Controller('database-migrations') export class DatabaseMigrationsController { constructor( private readonly databaseMigrationsService: DatabaseMigrationsService, ) {} private logger: Logger = new Logger(DatabaseMigrationsController.name); @Post('comment-test') async commentTest(): Promise<any> { try { await lastValueFrom(this.databaseMigrationsService.commentGenerationTest()); return []; } catch (error) { this.logger.error(error); throw new HttpException( { status: HttpStatus.INTERNAL_SERVER_ERROR, error: 'Comment generation failed', details: error, }, HttpStatus.INTERNAL_SERVER_ERROR ); } } }
database-migrations.service.ts
import { Injectable, Logger } from '@nestjs/common'; import fsp from 'fs/promises'; import path from 'path'; import { forkJoin, from, Observable, of, throwError } from 'rxjs'; import { catchError, map, switchMap, tap } from 'rxjs/operators'; import { AdditionalPropertiesService } from '../ai-core/additional-properties-predictor.service'; import { ColumnTestPredictorService, PredictedColumn as TestColumn, ValuePrediction, } from '../ai-core/column-test-predictor.service'; interface InputDataItem { textPrompt: string; table: { tableName: string; columns: LocalTestColumn[]; }; } interface LocalTestColumn { column_name: string; colConfidence: number; similarity: number; datatype: string; allowed_values: ValuePrediction; } @Injectable() export class DatabaseMigrationsService { private logger = new Logger(DatabaseMigrationsService.name); constructor( private readonly columnTestPredictorService: ColumnTestPredictorService, ) {} commentGenerationTest() { const tasks = [from(this.columnTestPredictorService.init())]; return forkJoin(tasks).pipe( switchMap(() => from( fsp.readFile(path.join(process.cwd(), 'src/input-data.json'), 'utf8'), ), ), switchMap((fileContent: string) => { this.logger.log('User propmts loaded'); const trainingData = JSON.parse( fileContent, ) as Partial<InputDataItem>[]; // Process each item in parallel return forkJoin( trainingData.map( (data: Partial<InputDataItem>) => forkJoin({ prediction: from( this.columnTestPredictorService.predictColumns( data.textPrompt as string, data.table?.tableName as string, data.table?.columns as Partial<TestColumn>[], ), ), }).pipe( map(({ prediction }) => ({ status: 'fulfilled', description: data.textPrompt, table: { tableName: data.table?.tableName, columns: prediction, }, })), ), catchError((error) => { this.logger.error(error); return of({ status: 'rejected', reason: error }); }), ), ).pipe( tap( ( results: { status: string; reason?: any; description?: string; table?: { columns: Partial<LocalTestColumn>[] }; }[], ) => { // Log rejected promises results .filter( (item: { status: string; reason?: any; description?: string; table?: { columns: Partial<LocalTestColumn>[] }; }) => item.status === 'rejected', ) .forEach( ( item: { status: string; reason?: any; description?: string; table?: { columns: Partial<LocalTestColumn>[] }; }, index: number, ) => { this.logger.error(`${index + 1}: ${item.reason}`); }, ); this.logger.debug('data', JSON.stringify(results, null, 2)); }, ), switchMap( ( results: { status: string; reason?: any; description?: string; table?: { columns: Partial<LocalTestColumn>[] }; }[], ) => from( fsp.writeFile( path.join(process.cwd(), 'src/whisperer.json'), JSON.stringify( results.filter( (item: { status: string; reason?: any; description?: string; table?: { columns: Partial<LocalTestColumn>[] }; }) => item.status === 'fulfilled', ), null, 2, ), ), ).pipe(map(() => results)), ), tap(() => this.logger.log('Whisperer finished')), catchError((error) => { this.logger.error(error); return throwError(error); }), ); }), map(() => []), ); } }
database.migrations.module.ts
import { Module } from '@nestjs/common'; import { ColumnTestPredictorService } from '../ai-core/column-test-predictor.service'; import { DatabaseMigrationsController } from './database-migrations.controller'; import { DatabaseMigrationsService } from './database-migrations.service'; @Module({ imports: [], controllers: [DatabaseMigrationsController], providers: [ DatabaseMigrationsController, DatabaseMigrationsService, ColumnTestPredictorService, ], exports: [ DatabaseMigrationsController, DatabaseMigrationsService, ColumnTestPredictorService, ], }) export class DatabaseMigrationsModule {}
app.module.ts
import { Module } from '@nestjs/common'; import { DatabaseMigrationsModule } from './ai/database-migrations/database-migrations.module'; @Module({ imports: [ DatabaseMigrationsModule, ], }) export class AppModule {}
main.ts
import compression from '@fastify/compress'; import { Logger } from '@nestjs/common'; import { NestFactory } from '@nestjs/core'; import { FastifyAdapter, NestFastifyApplication, } from '@nestjs/platform-fastify'; import { AppModule } from './app.module.js'; const logger = new Logger('main.ts'); async function bootstrap() { const app: NestFastifyApplication = await NestFactory.create<NestFastifyApplication>( AppModule, new FastifyAdapter(), { abortOnError: false, // Prevent NestJS from exiting on error }, ); app.setGlobalPrefix('/api'); app.enableCors(); await app.register(compression); await app.listen(3000); } bootstrap();
输入数据说明
输入数据包含多张数据库表的字段信息,结构包含表名、字段名、数据类型等内容。
输出结果说明
输出结果为生成的字段注释,多数注释不符合预期,仅小部分正确。
prompt.txt
Table name: team Column name: name Output: The name of the team Table name: team Column name: description Output: Description of the team Table name: team Column name: is_active Output: Whether the team is active Table name: grant Column name: title Output: Title of the grant Table name: grant Column name: amount Output: Funding amount Table name: grant Column name: start_date Output: Start date of the grant
项目文件结构
src ai ai-core column-test-predictor.ts database-migrations database-migrations.module.ts database-migrations.controller.ts database-migrations.service.ts models comment-predictor prompt.txt input-data.json whisperer.json app.module.ts main.ts
优化方案建议
1. 模型选型升级
flan-t5-small参数规模较小,对复杂语义理解能力有限。建议更换为更大的模型,比如Xenova/flan-t5-base或Xenova/flan-t5-large,这类模型具备更强的上下文理解和推理能力,能显著提升注释生成的准确性。
2. 优化提示词工程
- 补充更多样例:当前prompt仅包含基础样例,需添加更多覆盖不同场景的示例,比如包含外键关系、枚举字段、带单位的数值字段等场景,让模型学习到更全面的注释规则。
- 明确输出格式:在提示词中严格定义输出格式,比如要求注释必须以"表示..."或"存储..."开头,统一风格的同时减少模型输出的随机性。
- 加入数据类型信息:当前生成注释仅使用表名和字段名,可将字段的数据类型也加入提示词,比如
Table name: user, Column name: age, Datatype: INT, Output: 用户的年龄,帮助模型生成更精准的注释。
3. 调整生成参数
- 增大max_new_tokens:当前设置为10,部分合理注释可能被截断,建议调整为20-30,确保能生成完整的注释内容。
- 优化温度参数:temperature=0.3虽能降低随机性,但可能导致输出过于僵化。可尝试调整到0.5-0.7,在保证稳定性的同时保留一定灵活性;也可开启
top_p和top_k参数,比如设置top_p: 0.9, top_k: 50,让模型从更优质的候选词中选择。
4. 加入上下文信息
如果表之间存在关联关系,可将表的业务含义或关联表信息加入提示词,比如Table name: order_item, Column name: order_id, Related table: order, Output: 关联订单的ID,帮助模型理解字段的业务逻辑。
5. 批量处理优化
当前采用串行循环生成注释,可改为批量处理,将多个字段的请求合并为一个prompt,让模型同时处理多个字段,利用上下文关联提升注释的一致性和准确性。
6. 结果校验与修正
添加后处理逻辑,对生成的注释进行校验:
- 过滤无意义的输出,比如仅包含字段名的注释;
- 对不符合格式要求的注释进行修正,比如统一术语(将"是否活跃"改为"是否启用");
- 针对常见字段(如id、create_time)直接使用预设注释,无需模型生成。
内容的提问来源于stack exchange,提问作者Orist Timemaker
相关产品推荐
相关产品推荐

