如何用SQL/SQLAlchemy生成符合规则的不重复组合并写入Table1?
基于SQL与SQLAlchemy的组合生成方案
SQL实现(以PostgreSQL为例)
1. 提取分组数据
先从Table2中聚合各分组的有效值,再拆分为单独数据集用于生成组合:
WITH group_values AS ( SELECT group_num, ARRAY_AGG(value) AS vals FROM Table2 WHERE group_num IN (1,2,3) GROUP BY group_num ), group1 AS (SELECT unnest(vals) AS val FROM group_values WHERE group_num=1), group2 AS (SELECT unnest(vals) AS val FROM group_values WHERE group_num=2), group3 AS (SELECT unnest(vals) AS val FROM group_values WHERE group_num=3)
2. 生成去重组合
通过对组合元素排序生成唯一标识,过滤掉a,b,c和b,a,c这类重复组合:
, unique_combinations AS ( SELECT DISTINCT ARRAY_TO_STRING(ARRAY[g1.val, g2.val, g3.val] ORDER BY g1.val, g2.val, g3.val, ',') AS combo_key, g1.val AS val1, g2.val AS val2, g3.val AS val3 FROM group1 g1 CROSS JOIN group2 g2 CROSS JOIN group3 g3 )
3. 过滤互斥组合
检查Table1中已存在的g类型组合,排除掉规则中互斥的组合(需根据实际数值与字符的映射关系调整REPLACE逻辑):
, filtered_combinations AS ( SELECT uc.* FROM unique_combinations uc LEFT JOIN Table1 t1 ON t1.type = 'g' AND ( -- 匹配对应数值组合 (t1.val1 = REPLACE(uc.val1, 'a', '1') AND t1.val2 = REPLACE(uc.val2, 'b', '1') AND t1.val3 = REPLACE(uc.val3, 'c', '2')) -- 匹配对应字符组合 OR (t1.val1 = REPLACE(uc.val1, '1', 'a') AND t1.val2 = REPLACE(uc.val2, '1', 'b') AND t1.val3 = REPLACE(uc.val3, '2', 'c')) ) WHERE t1.id IS NULL )
4. 插入到Table1
将符合规则的组合写入Table1:
INSERT INTO Table1 (type, val1, val2, val3) SELECT 'g', val1, val2, val3 FROM filtered_combinations;
SQLAlchemy实现
1. 定义数据表模型
from sqlalchemy import Column, Integer, String from sqlalchemy.ext.declarative import declarative_base Base = declarative_base() class Table1(Base): __tablename__ = 'table1' id = Column(Integer, primary_key=True) type = Column(String) val1 = Column(String) val2 = Column(String) val3 = Column(String) class Table2(Base): __tablename__ = 'table2' id = Column(Integer, primary_key=True) group_num = Column(Integer) value = Column(String)
2. 获取分组有效值
from sqlalchemy import create_engine from sqlalchemy.orm import sessionmaker # 初始化数据库连接 engine = create_engine('your_database_url') Session = sessionmaker(bind=engine) session = Session() # 提取各分组的值 group1_vals = [row[0] for row in session.query(Table2.value).filter(Table2.group_num == 1).all()] group2_vals = [row[0] for row in session.query(Table2.value).filter(Table2.group_num == 2).all()] group3_vals = [row[0] for row in session.query(Table2.value).filter(Table2.group_num == 3).all()]
3. 生成去重组合
用itertools生成笛卡尔积,通过排序去重:
import itertools # 生成所有可能的组合 all_combos = list(itertools.product(group1_vals, group2_vals, group3_vals)) # 去重:保留排序后唯一的组合 seen = set() unique_combos = [] for combo in all_combos: sorted_combo = tuple(sorted(combo)) if sorted_combo not in seen: seen.add(sorted_combo) unique_combos.append(combo)
4. 过滤互斥组合
# 获取已存在的g类型组合 existing_combos = set(session.query(Table1.val1, Table1.val2, Table1.val3).filter(Table1.type == 'g').all()) # 定义数值与字符的互斥映射(需根据实际业务调整) mutex_map = {'1':'a', 'a':'1', '2':'c', 'c':'2', 'b':'1', '1':'b'} # 过滤掉互斥组合 filtered_combos = [] for combo in unique_combos: mutex_combo = tuple(mutex_map.get(val, val) for val in combo) if combo not in existing_combos and mutex_combo not in existing_combos: filtered_combos.append(combo)
5. 批量插入数据
# 创建新记录并批量插入 new_records = [Table1(type='g', val1=v1, val2=v2, val3=v3) for v1, v2, v3 in filtered_combos] session.add_all(new_records) session.commit() session.close()
内容的提问来源于stack exchange,提问作者devb
相关产品推荐
相关产品推荐

