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

如何为SDV HMASynthesizer添加关系自定义逻辑并提升分类数据质量?

问题描述

我正在使用SDV的HMASynthesizer生成合成数据,遇到两个核心问题:

  1. 表间关系自定义逻辑失效:users表中role为mentor的user_id只能出现在sessions表的mentor_id列,role为mentee的user_id只能出现在mentee_id列,但生成的数据出现角色不匹配(比如user_id=2566的用户是mentee,却出现在sessions的mentor_id列)。
  2. 分类数据生成质量待提升:当前使用LabelEncoder(add_noise=False)处理分类字段,希望优化分类数据的生成效果。

完整模型代码

database_data = {
    'domain': domain,
    'region': region,
    'sessions': sessions,
    'users': users
}

database_metadata = MultiTableMetadata()

for i in database_data:
    database_metadata.detect_table_from_dataframe(
        table_name=i,
        data=database_data[i]
    )
# sessions________________________

database_metadata.update_column(
    table_name='sessions',
    column_name='session_id',
    sdtype='id',
    regex_format='[0-9]{5}'
)

database_metadata.update_column(
    table_name='sessions',
    column_name='mentor_id',
    sdtype='id',
    regex_format='[a-zA-Z]{4}'
)

database_metadata.update_column(
    table_name='sessions',
    column_name='mentee_id',
    sdtype='id',
    regex_format='[a-zA-Z]{4}'
)

database_metadata.update_column(
    table_name='sessions',
    column_name='mentor_domain_id',
    sdtype='id',
    regex_format='[a-zA-Z]{2}'
)

database_metadata.set_primary_key(
    table_name='sessions',
    column_name='session_id'
)

# users________________________

database_metadata.update_column(
    table_name='users',
    column_name='user_id',
    sdtype='id',
    regex_format='[0-9]{4}'
)


database_metadata.update_column(
    table_name='users',
    column_name='region_id',
    sdtype='id',
    regex_format='[a-zA-Z]{2}'
)

database_metadata.set_primary_key(
    table_name='users',
    column_name='user_id'
)

# domain ________________________

database_metadata.update_column(
    table_name='domain',
    column_name='id',
    sdtype='id',
    regex_format='[0-9]{2}'
)

database_metadata.set_primary_key(
    table_name='domain',
    column_name='id'
)

# region _______________________

database_metadata.update_column(
    table_name='region',
    column_name='id',
    sdtype='id',
    regex_format='[0-9]{2}'
)

database_metadata.set_primary_key(
    table_name='region',
    column_name='id'
)

# add relationship

database_metadata.add_relationship(
    parent_table_name='domain',
    child_table_name='sessions',
    parent_primary_key='id',
    child_foreign_key='mentor_domain_id'
)

database_metadata.add_relationship(
    parent_table_name='region',
    child_table_name='users',
    parent_primary_key='id',
    child_foreign_key='region_id'
)

database_metadata.add_relationship(
    parent_table_name='users',
    child_table_name='sessions',
    parent_primary_key='user_id',
    child_foreign_key='mentor_id'
)

database_metadata.add_relationship(
    parent_table_name='users',
    child_table_name='sessions',
    parent_primary_key='user_id',
    child_foreign_key='mentee_id'
)

database_metadata.visualize(
    show_table_details=True,
    show_relationship_labels=True,
    output_filepath='my_metadata.png'
)

# Synthesizer
synthesizer = HMASynthesizer(database_metadata, locales=['ru_RU'])

# transformers
synthesizer.auto_assign_transformers(database_data)

from rdt.transformers.categorical import LabelEncoder

synthesizer.update_transformers(
    table_name='domain',
    column_name_to_transformer={
        'name': LabelEncoder(add_noise=False)
    }
)

synthesizer.update_transformers(
    table_name='region',
    column_name_to_transformer={
        'name': LabelEncoder(add_noise=False)
    }
)

synthesizer.update_transformers(
    table_name='users',
    column_name_to_transformer={
        'role': LabelEncoder(add_noise=False)
    }
)

synthesizer.update_transformers(
    table_name='sessions',
    column_name_to_transformer={
        'session_status': LabelEncoder(add_noise=False)
    }
)

# preprocess data
processed_data = synthesizer.preprocess(database_data)

# model fit
synthesizer.fit_processed_data(processed_data)

synthesizer.reset_sampling()
database_syntetic_data = synthesizer.sample(scale=1.01)

错误示例

users表错误记录

user_idreg_dateroleregion_id
255625562022-08-13mentee17

sessions表错误记录

session_idsession_date_timementor_idmentee_idsession_statusmentor_domain_id
1701702021-08-0425661720finished0
129612962021-12-172566431canceled1
449744972022-05-1625661327canceled4
542954292022-03-1125661543canceled5

解决方案

问题1:实现用户角色与会话字段的关联约束

SDV默认外键仅保证引用完整性,不会关联role逻辑,可通过以下两种方式解决:

方案1:采样后数据清洗(快速落地)

生成数据后直接修正不符合规则的记录:

import numpy as np

# 提取合成数据中的用户和会话表
users_synth = database_syntetic_data['users']
sessions_synth = database_syntetic_data['sessions']

# 分离合法的mentor和mentee ID集合
valid_mentors = set(users_synth[users_synth['role'] == 'mentor']['user_id'])
valid_mentees = set(users_synth[users_synth['role'] == 'mentee']['user_id'])

# 修正mentor_id列:替换非法ID为随机合法mentor ID
invalid_mentor_rows = ~sessions_synth['mentor_id'].isin(valid_mentors)
sessions_synth.loc[invalid_mentor_rows, 'mentor_id'] = np.random.choice(
    list(valid_mentors),
    size=invalid_mentor_rows.sum()
)

# 修正mentee_id列:替换非法ID为随机合法mentee ID
invalid_mentee_rows = ~sessions_synth['mentee_id'].isin(valid_mentees)
sessions_synth.loc[invalid_mentee_rows, 'mentee_id'] = np.random.choice(
    list(valid_mentees),
    size=invalid_mentee_rows.sum()
)

# 更新合成数据集
database_syntetic_data['sessions'] = sessions_synth

方案2:模型层面添加自定义约束(从根源避免错误)

通过SDV的CustomConstraint定义规则,让模型生成时自动遵守:

from sdv.constraints import CustomConstraint

def validate_session_role_assignment(sessions_table, users_table):
    # 验证mentor_id对应的用户角色为mentor
    mentor_check = sessions_table.merge(
        users_table,
        left_on='mentor_id',
        right_on='user_id',
        how='left'
    )['role'] == 'mentor'
    
    # 验证mentee_id对应的用户角色为mentee
    mentee_check = sessions_table.merge(
        users_table,
        left_on='mentee_id',
        right_on='user_id',
        how='left'
    )['role'] == 'mentee'
    
    # 返回所有合法的行
    return mentor_check & mentee_check

# 创建约束实例
role_link_constraint = CustomConstraint(
    validate_fn=validate_session_role_assignment,
    tables=['sessions', 'users']
)

# 初始化合成器时传入约束
synthesizer = HMASynthesizer(
    database_metadata,
    locales=['ru_RU'],
    constraints=[role_link_constraint]
)

注意:自定义约束会增加训练和采样时间,数据量大时需权衡效率。

问题2:提升分类数据生成质量

LabelEncoder对无序分类的学习能力有限,推荐以下优化方向:

1. 替换为更适配的编码器

  • OneHotEncoder:适合无序分类,保留类别间独立性:
from rdt.transformers.categorical import OneHotEncoder

# 替换users表role字段的编码器
synthesizer.update_transformers(
    table_name='users',
    column_name_to_transformer={
        'role': OneHotEncoder()
    }
)
  • CategoryEncoder:SDV专属编码器,更擅长学习复杂分类分布:
from sdv.transformers.categorical import CategoryEncoder

# 替换domain和region表name字段的编码器
synthesizer.update_transformers(
    table_name='domain',
    column_name_to_transformer={
        'name': CategoryEncoder()
    }
)

synthesizer.update_transformers(
    table_name='region',
    column_name_to_transformer={
        'name': CategoryEncoder()
    }
)

2. 调整模型训练参数

增加训练轮数,让模型更充分学习分类特征分布:

synthesizer = HMASynthesizer(
    database_metadata,
    locales=['ru_RU'],
    epochs=100,  # 默认50,可根据数据规模调整至50-200
    verbose=True  # 查看训练进度
)

3. 量化评估分类数据质量

使用SDV的评估工具验证生成效果,迭代优化:

from sdv.evaluation.multi_table import evaluate_quality, get_column_plot

# 生成质量报告
quality_report = evaluate_quality(
    real_data=database_data,
    synthetic_data=database_syntetic_data,
    metadata=database_metadata
)

# 查看分类字段的分布对比
get_column_plot(
    real_data=database_data['users'],
    synthetic_data=database_syntetic_data['users'],
    column_name='role',
    plot_type='bar'
)

内容的提问来源于stack exchange,提问作者John Doe

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 05:37:03