统计推断 统计推断(statistical inference)不止于给出是/否的决定,而是带着可量化的不确定性去估计总体参数。本文件涵盖置信区间、点估计与区间估计、极大似然估计、矩估计法以及回归分析,这是连接原始数据与 ML 预测模型之间的桥梁。 假设检验给你一个是/否的决定:拒绝或不拒绝。但很多时候你想要更有信息量的东西——给你所估计的参数一个合理的取值范围。这正是置信区间(confidence intervals)所提供的。 点估计(point estimate)是从样本中算出的一个单一数字,比如样本均值 $\bar{x}$。它是你对总体参数的最佳猜测,但它本身并不能告诉你这个估计有多精确。 置信区间用一段范围把点估计包起来,以反映不确定性。
统计推断(statistical inference)不止于给出是/否的决定,而是带着可量化的不确定性去估计总体参数。本文件涵盖置信区间、点估计与区间估计、极大似然估计、矩估计法以及回归分析,这是连接原始数据与 ML 预测模型之间的桥梁。
假设检验给你一个是/否的决定:拒绝或不拒绝。但很多时候你想要更有信息量的东西——给你所估计的参数一个合理的取值范围。这正是**置信区间(confidence intervals)**所提供的。
**点估计(point estimate)**是从样本中算出的一个单一数字,比如样本均值 \bar{x}。它是你对总体参数的最佳猜测,但它本身并不能告诉你这个估计有多精确。
置信区间用一段范围把点估计包起来,以反映不确定性。它的形式是:
95% 置信区间的意思是:如果你把实验重复很多次,每次都构造一个区间,那么这些区间里大约有 95% 会包含真实的总体参数。它并不意味着「这个特定区间包含参数的概率是 95%」。参数是固定的,变化的是那些区间。
例题:你测量了 50 个人的身高,得到 \bar{x} = 170 cm,\sigma = 8 cm。构造一个 95% 置信区间。
你可以有 95% 的把握说,真实的平均身高落在 167.78 到 172.22 cm 之间。
当 \sigma 未知时(这是常见情形),改用样本标准差 s 和 t 分布:
区间越宽越有把握,但越不精确。区间越窄越精确,但把握越小。要想在不损失置信度的前提下让区间变窄,唯一的办法是增大样本量。
**功效分析(power analysis)**帮助你在动手做实验之前先做好规划。问题是:为了以指定的功效察觉到某个给定大小的效应,我需要多大的样本量?
回忆上一个文件,功效 = 1 - \beta,即正确拒绝一个为假的 H_0 的概率。一个常见目标是 80% 的功效。
用 z 检验去察觉差异 \delta、显著性为 \alpha、功效为 1-\beta 时,所需的样本量是:
你大约需要每组 126 人。
功效分析可以避免两种常见错误:把实验做得太小以至于察觉不到真实效应(功效不足,underpowered),或者把实验做得远超所需而浪费资源(功效过剩,overpowered)。
**蒙特卡洛方法(Monte Carlo methods)**用随机抽样来解决那些难以或无法用解析方法求解的问题。核心思路是:如果某个东西你算不准,那就模拟它很多次,把结果当作近似值。
这个名字来源于蒙特卡洛赌场,是对随机性所扮演角色的致意。这些方法是 ML 中的主力工具,用于估计积分、评估模型不确定性以及逼近复杂分布等任务。
通用的蒙特卡洛配方:
一个经典的例子是估计 \pi。想象一个边长为 2、以原点为中心的正方形,里面内切一个半径为 1 的圆。正方形面积是 4,圆面积是 \pi。
一个点 (x, y) 落在圆内的条件是 x^2 + y^2 \le 1。投的点越多,你的估计就越接近 \pi 的真实值。
在 ML 中,蒙特卡洛方法出现在:
**因子分析(factor analysis)**是一种用来发现隐藏(潜)变量、并用它们解释可观测变量之间相关性的技术。如果 10 道性格调查题可以被 3 个潜在特质(外向性、宜人性、尽责性)所解释,因子分析就能找出这些特质。
该模型假设每个可观测变量 x_i 都是少数几个潜因子 f_j 的线性组合再加上噪声:
这些 \lambda 值称为因子载荷(factor loadings),告诉你每个可观测变量与每个因子之间关联有多强。这与第 2 章的矩阵分解直接呼应;因子分析与特征值分解和 SVD 关系密切。
**实验设计(experimental design)**是安排实验结构、使你能够得出有效结论的艺术。糟糕的设计甚至能让一个庞大的数据集变得毫无用处。
一个设计良好的实验包含的关键要素:
常见的实验设计:
在 ML 实验中,这些原则至关重要。比较模型时,你应该控制随机种子、数据集划分和硬件。交叉验证就是交叉设计的一种形式。消融实验——每次移除一个组件——遵循的正是析因设计的逻辑。
import jax.numpy as jnp x_bar = 170.0 # 样本均值 sigma = 8.0 # 总体标准差(已知) n = 50 # 样本量 # 常见置信水平对应的临界值 z_stars = {0.90: 1.645, 0.95: 1.960, 0.99: 2.576} for conf, z_star in z_stars.items(): me = z_star * (sigma / jnp.sqrt(n)) lower, upper = x_bar - me, x_bar + me print(f"{conf*100:.0f}% CI: [{lower:.2f}, {upper:.2f}] (ME = {me:.2f})")
import jax import jax.numpy as jnp import matplotlib.pyplot as plt key = jax.random.PRNGKey(42) # 在 [-1, 1] x [-1, 1] 中生成随机点 n_points = 100_000 k1, k2 = jax.random.split(key) x = jax.random.uniform(k1, shape=(n_points,), minval=-1, maxval=1) y = jax.random.uniform(k2, shape=(n_points,), minval=-1, maxval=1) # 检查哪些点落在单位圆内 inside = (x**2 + y**2) <= 1.0 cumulative_inside = jnp.cumsum(inside) counts = jnp.arange(1, n_points + 1) pi_estimates = 4.0 * cumulative_inside / counts plt.figure(figsize=(10, 4)) plt.plot(pi_estimates, color="#3498db", alpha=0.7, linewidth=0.5) plt.axhline(y=jnp.pi, color="#e74c3c", linestyle="--", label=f"π = {jnp.pi:.6f}") plt.xlabel("Number of points") plt.ylabel("Estimate of π") plt.title("Monte Carlo estimation of π") plt.legend() plt.ylim(2.8, 3.5) plt.show() print(f"Final estimate: {pi_estimates[-1]:.6f}") print(f"True value: {jnp.pi:.6f}") print(f"Error: {abs(pi_estimates[-1] - jnp.pi):.6f}")
import jax import jax.numpy as jnp # 参数 delta = 2.0 # 效应量(均值之差) sigma = 8.0 # 总体标准差 alpha = 0.05 power_target = 0.80 # 解析地求样本量 z_alpha = 1.96 # 双尾,alpha=0.05 z_beta = 0.84 # 功效=0.80 n_required = ((z_alpha + z_beta) * sigma / delta) ** 2 print(f"Required n per group: {n_required:.0f}") # 用模拟验证 key = jax.random.PRNGKey(7) n = int(jnp.ceil(n_required)) n_sims = 5000 rejections = 0 for _ in range(n_sims): key, k1, k2 = jax.random.split(key, 3) group_a = jax.random.normal(k1, shape=(n,)) * sigma + 50 group_b = jax.random.normal(k2, shape=(n,)) * sigma + 50 + delta pooled_se = jnp.sqrt(2 * sigma**2 / n) z = (group_b.mean() - group_a.mean()) / pooled_se p = 2 * (1 - __import__("jax").scipy.stats.norm.cdf(jnp.abs(z))) if p <= alpha: rejections += 1 print(f"Simulated power: {rejections/n_sims:.3f}") print(f"Target power: {power_target:.3f}")
import jax.numpy as jnp import matplotlib.pyplot as plt sigma = 8.0 z_star = 1.96 # 95% 置信 sample_sizes = jnp.array([10, 20, 30, 50, 100, 200, 500, 1000], dtype=jnp.float32) margins = z_star * sigma / jnp.sqrt(sample_sizes) plt.figure(figsize=(8, 4)) plt.bar([str(int(n)) for n in sample_sizes], margins, color="#3498db", alpha=0.7) plt.xlabel("Sample size") plt.ylabel("Margin of error (cm)") plt.title("95% CI margin of error shrinks with larger samples") plt.show()