如何将计算硬币序列条件概率的R代码转为Netezza SQL?
问题描述
我有如下R代码,用于模拟不同学生抛硬币的随机数据,统计所有学生得到的序列组合数,并计算条件概率(已知前两次结果时第三次结果的概率):
library(dplyr) library(tidyverse) ids = 1:100 student_id = sample(ids, 1000, replace = TRUE) coin_result = sample(c("H", "T"), 1000, replace = TRUE) my_data = data.frame(student_id, coin_result) my_data = my_data[order(my_data$student_id),] my_data = my_data %>% group_by(student_id) %>% summarize(Sequence = str_c(coin_result, lead(coin_result), lead(coin_result, 2)), .groups = 'drop') %>% filter(!is.na(Sequence)) %>% count(Sequence) final = my_data %>% mutate(two_seq = substr(Sequence, 1, 2)) %>% group_by(two_seq) %>% mutate(third = substr(Sequence, 3, 3)) %>% group_by(two_seq, third) %>% summarize(sums = sum(n)) %>% mutate(prob = sums / sum(sums)) # FINAL RESULT (Note: I think this is correct because the pairs of conditional probabilities add to 1) # A tibble: 8 x 4 # Groups: two_seq [4] two_seq third sums prob <chr> <chr> <int> <dbl> 1 HH H 112 0.514 2 HH T 106 0.486 3 HT H 108 0.537 4 HT T 93 0.463 5 TH H 97 0.5 6 TH T 97 0.5 7 TT H 93 0.497 8 TT T 94 0.503
我正尝试将上述R代码转换为(Netezza)SQL。我查阅资料后了解到LEAD()函数可实现该功能,以下是我的尝试:
SELECT Sequence, COUNT(*) as count FROM ( SELECT STUDENT_ID, coin_result || LEAD(coin_result, 1) OVER (PARTITION BY STUDENT_ID ORDER BY STUDENT_ID) || LEAD(coin_result, 2) OVER (PARTITION BY STUDENT_ID ORDER BY STUDENT_ID) || LEAD(coin_result, 3) OVER (PARTITION BY STUDENT_ID ORDER BY STUDENT_ID) AS Sequence FROM my_data ) subquery WHERE Sequence NOT LIKE '%NULL%' GROUP BY Sequence;
我不确定当前写法是否正确,也不知道如何添加条件概率列,恳请指导我完成后续步骤。谢谢!
解决方案
第一步:修正序列生成的SQL代码
你的现有SQL存在两个核心问题:
- R代码生成的是3个连续结果的序列(当前结果+后1位+后2位),但你的SQL拼接了4个值,和原逻辑不符
- 分区内仅按
STUDENT_ID排序无意义,需要给每个学生的抛硬币记录生成明确的顺序,否则LEAD无法取到连续的结果
修正后的序列生成SQL:
WITH student_coin_orders AS ( SELECT STUDENT_ID, coin_result, -- 给每个学生的抛硬币记录生成顺序号,确保LEAD取到连续结果 ROW_NUMBER() OVER(PARTITION BY STUDENT_ID ORDER BY CURRENT_TIMESTAMP) AS coin_order FROM my_data ), three_sequences AS ( SELECT STUDENT_ID, -- 拼接当前+后1+后2位,生成3位序列 coin_result || LEAD(coin_result, 1) OVER(PARTITION BY STUDENT_ID ORDER BY coin_order) || LEAD(coin_result, 2) OVER(PARTITION BY STUDENT_ID ORDER BY coin_order) AS Sequence FROM student_coin_orders ) SELECT Sequence, COUNT(*) AS n FROM three_sequences -- 过滤掉不完整的序列(每个学生最后2条记录会生成含NULL的序列) WHERE Sequence IS NOT NULL GROUP BY Sequence;
第二步:计算条件概率
基于上述序列统计结果,使用窗口函数SUM() OVER()计算每个前两位序列的总次数,进而得到条件概率:
WITH student_coin_orders AS ( SELECT STUDENT_ID, coin_result, ROW_NUMBER() OVER(PARTITION BY STUDENT_ID ORDER BY CURRENT_TIMESTAMP) AS coin_order FROM my_data ), three_sequences AS ( SELECT STUDENT_ID, coin_result || LEAD(coin_result, 1) OVER(PARTITION BY STUDENT_ID ORDER BY coin_order) || LEAD(coin_result, 2) OVER(PARTITION BY STUDENT_ID ORDER BY coin_order) AS Sequence FROM student_coin_orders ), sequence_counts AS ( SELECT Sequence, COUNT(*) AS n FROM three_sequences WHERE Sequence IS NOT NULL GROUP BY Sequence ) SELECT -- 提取前两位序列 SUBSTR(Sequence, 1, 2) AS two_seq, -- 提取第三位结果 SUBSTR(Sequence, 3, 3) AS third, SUM(n) AS sums, -- 计算条件概率:当前组合次数 / 前两位序列的总次数,保留3位小数 ROUND(SUM(n) * 1.0 / SUM(SUM(n)) OVER(PARTITION BY SUBSTR(Sequence, 1, 2)), 3) AS prob FROM sequence_counts GROUP BY SUBSTR(Sequence, 1, 2), SUBSTR(Sequence, 3, 3) ORDER BY two_seq, third;
关键说明
- Netezza中字符串拼接用
||是正确的,对应R中的str_c - 必须用
ROW_NUMBER()生成顺序号,否则LEAD的结果无序,会导致序列错误 SUM(SUM(n)) OVER(PARTITION BY two_seq)是Netezza支持的窗口函数写法,比子查询计算分组总和效率更高ROUND(..., 3)用来和R结果的小数位数对齐,可根据需求调整
内容的提问来源于stack exchange,提问作者stats_noob
相关产品推荐
相关产品推荐

