"""Pearson、Spearman与控制协变量后的偏相关。""" import numpy as np from scipy import stats from flask_babel import gettext as _ from .contracts import Column, PaperTable, StatisticalResult, columns, numeric_frame, enough, varying, choice, boolean def correlation(data, options, method): selected = columns(data, options.get('variables'), 2, 50) controls = [] if method == 'stat_partial_correlation': controls = columns(data, options.get('controls')) if set(selected).intersection(controls): raise ValueError(_('控制变量不能与分析变量重复')) policy = choice(options, 'missing', 'complete', ('complete', 'pairwise')) if controls and policy != 'complete': raise ValueError(_('偏相关使用统一完整样本,请选择完整案例')) numeric = numeric_frame(data, selected + controls) if policy == 'complete': numeric = enough(numeric.dropna(), len(controls)+3) residuals = None if controls: design = np.column_stack([np.ones(len(numeric)), numeric[controls].values]) if np.linalg.matrix_rank(design) < design.shape[1]: raise ValueError(_('控制变量存在完全共线性,请减少重复信息的变量')) residuals = numeric[selected].values - design @ np.linalg.lstsq(design, numeric[selected].values, rcond=None)[0] rows, pairs = [], [] for i, name in enumerate(selected): row = {'variable': name, 'n': int(numeric[name].notna().sum()), 'mean': numeric[name].mean(), 'sd': numeric[name].std(ddof=1)} for j, other in enumerate(selected): if j > i: row['r'+str(j)] = None continue sample = enough(numeric[[name, other]].dropna() if i != j else numeric[[name]].dropna(), 3) if i == j: varying(sample[name].values, name) row['r'+str(j)] = (1.0, None) continue if residuals is not None: x, y = residuals[:, i], residuals[:, j] else: x, y = sample[name].values, sample[other].values varying(x, name) varying(y, other) if method == 'stat_spearman': r, p = stats.spearmanr(x, y) else: r, p = stats.pearsonr(x, y) degrees = len(x)-len(controls)-2 if controls: p = 0.0 if abs(r) >= 1 else 2*stats.t.sf(abs(r)*np.sqrt(degrees/(1-r*r)), degrees) row['r'+str(j)] = (float(r), float(p)) pairs.append({'left': name, 'right': other, 'n': len(x), 'r': r, 'p': p, 'df': degrees}) rows.append(row) headers = [Column('variable', _('变量'), 'text')] if boolean(options, 'include_descriptives', True): headers += [Column('n', 'N', 'integer'), Column('mean', _('均值')), Column('sd', _('标准差'))] headers += [Column('r'+str(i), str(i+1), 'correlation') for i in range(len(selected))] for i, row in enumerate(rows): row['variable'] = str(i+1)+'. '+row['variable'] return StatisticalResult([PaperTable(_('相关分析'), headers, rows, _('* p<0.05,** p<0.01,*** p<0.001。'))], {'pairs': pairs, 'controls': controls, 'missing': policy, 'n_by_variable': {name: int(numeric[name].notna().sum()) for name in selected}})