如何为SDV HMASynthesizer添加关系自定义逻辑并提升分类数据质量?
问题描述
我正在使用SDV的HMASynthesizer生成合成数据,遇到两个核心问题:
- 表间关系自定义逻辑失效:
users表中role为mentor的user_id只能出现在sessions表的mentor_id列,role为mentee的user_id只能出现在mentee_id列,但生成的数据出现角色不匹配(比如user_id=2566的用户是mentee,却出现在sessions的mentor_id列)。 - 分类数据生成质量待提升:当前使用
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_id | reg_date | role | region_id | |
|---|---|---|---|---|
| 2556 | 2556 | 2022-08-13 | mentee | 17 |
sessions表错误记录
| session_id | session_date_time | mentor_id | mentee_id | session_status | mentor_domain_id | |
|---|---|---|---|---|---|---|
| 170 | 170 | 2021-08-04 | 2566 | 1720 | finished | 0 |
| 1296 | 1296 | 2021-12-17 | 2566 | 431 | canceled | 1 |
| 4497 | 4497 | 2022-05-16 | 2566 | 1327 | canceled | 4 |
| 5429 | 5429 | 2022-03-11 | 2566 | 1543 | canceled | 5 |
解决方案
问题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
相关产品推荐
相关产品推荐

