You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.12 22:27:02