1. 生存分析里的建模目标:为什么偏要造一个"Cox loss"
我最早接触 Cox loss,是在做客户流失预测的时候。当时团队要做的不只是"判断哪些用户会走",而是"预测用户多久之后会走",因为续费提醒、挽留策略的投放时机,全依赖那个时间点。试了一圈常规分类模型,发现都不对劲:训练数据里有一大批用户,观察期结束还没流失,你没法给他们标一个确切的流失时间,但直接丢掉又太浪费。查了一圈资料,才发现这类问题有个正经名字叫生存分析,而 Cox 比例风险模型就是其中最常用的武器,所谓的 Cox loss,正是训练这个模型时用的损失函数。
先说清楚一个容易混淆的点:Cox loss 不是一个"损失函数"的名字,而是Cox 比例风险模型采用负对数偏似然(negative log partial likelihood)作为优化目标时的叫法。在传统统计软件里,你可能没见过"loss"这种说法,大家习惯叫"偏似然最大化";但到了深度学习时代,所有东西都要变成可微分的损失函数,Cox loss 这个叫法就自然流传开了。
它的核心用途,是拟合下面这个风险函数:
$$h(t|x) = h_0(t) \cdot \exp(f(x))$$
其中 $h(t|x)$ 表示特征为 $x$ 的个体在时刻 $t$ 的瞬时风险,$h_0(t)$ 是基线风险函数,与特征无关,只随时间变化;而 $f(x)$ 是我们真正关心的部分——它把原始特征映射成一个线性组合或者神经网络输出。Cox 模型的巧妙之处在于,它不关心 $h_0(t)$ 长什么样,只估计 $f(x)$ 那一部分,所以它属于半参数模型。
这种设计带来的直接好处是:你不需要对生存时间分布做任何假设(不需要猜它是指数分布还是 Weibull 分布),只需要假设不同个体之间的风险是成比例的,也就是任意两个个体 $i$ 和 $j$,在任意时刻 $t$,它们的风险比 $\frac{h(t|x_i)}{h(t|x_j)}$ 是常数,不随时间变化。这个假设是所有 Cox 模型应用的前提,模型效果崩了的时候,第一个要怀疑的就是它。
那损失函数怎么构造呢?核心思路是:我们不知道基线风险,但可以绕过它,通过比较同一时刻"谁先出事、谁还没出事"来估计参数。这就是偏似然的核心思想——只用"事件发生的顺序"来拟合,不需要精确的时间分布。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 偏似然损失公式拆解:从风险集到负对数
Cox loss 的公式看起来不长,但每个符号背后都有具体的统计含义。先把完整的负对数偏似然写出来:
$$L(\theta) = -\sum_{i: \delta_i = 1} \left[ f(x_i) - \log \sum_{j \in R(t_i)} \exp(f(x_j)) \right]$$
看着有点吓人,拆开就清楚了。这里的符号约定是:
- $x_i$:第 $i$ 个样本的特征向量
- $t_i$:第 $i$ 个样本的观测时间(可能是事件发生时间,也可能是删失时间)
- $\delta_i$:事件指示符,$\delta_i = 1$ 表示观测到了事件发生,$\delta_i = 0$ 表示右删失(只观测到"到某个时间还没发生")
- $f(x_i)$:模型输出的风险评分(log-risk),可以是线性回归的结果,也可以是神经网络最后一层的输出
- $R(t_i)$:风险集(risk set),指在时间 $t_i$ 时刻仍然"存活"(未被删失、未发生事件)的所有样本的集合
2.1 为什么只看"发生事件"的样本
注意求和符号下面的条件 $\delta_i = 1$——只有真正发生了事件的样本才会进入损失计算。删失样本不直接参与分子,但它们会出现在分母的风险集里。
举个例子,假设有 5 个病人,观察期内 3 人死亡、2 人失访(删失)。第一个人在第 10 天死亡,那么在他死亡的那一刻,风险集里包含所有在第 10 天还没死的病人——包括后来在第 20 天、第 30 天死亡的那两位,也包括到最后都没死的删失样本。这个时刻的信息是:"第 10 天,有一堆人还活着,但偏偏是这个人死了,为什么?"
这就是 Cox 模型的建模逻辑——它不预测绝对风险,而是预测相对风险排序。谁会在哪个时刻出事不重要,重要的是"在某个时刻,出事的是不是你"。
2.2 分子分母的含义:个体风险 vs 群体风险
分子 $\exp(f(x_i))$ 表示第 $i$ 个样本在当前时刻的"风险强度";分母 $\sum_{j \in R(t_i)} \exp(f(x_j))$ 表示风险集里所有样本的风险强度之和。两者相除,得到的是"在 $t_i$ 时刻,第 $i$ 个样本出事在所有可能出事的人当中所占的比例"。
模型的优化目标很直接:让这个比例尽量大。也就是说,发生在 $t_i$ 时刻的事件,我们希望模型预测"就是这个人最该出事",而不是别人。把所有事件时刻的这个比例乘起来,就是偏似然函数;取负对数,就变成求和形式的损失,方便梯度下降。
2.3 为什么取负对数
取负对数有三个实际理由。第一,事件概率是大量小数连乘,数值下溢到零是家常便饭,取对数能把乘法变加法;第二,优化器习惯最小化损失,加个负号把最大化问题转成最小化;第三,对数里恰好出现"log-sum-exp"结构,这在数值计算上有稳定的实现方式。
如果你手推一下梯度就会发现,这个损失函数的梯度有一个很优雅的解释:损失对 $f(x_i)$ 的梯度,等于"预测的风险比例"减去"实际的事件指示":
$$\frac{\partial L}{\partial f(x_i)} = \frac{\exp(f(x_i))}{\sum_{j \in R(t_i)} \exp(f(x_j))} - \delta_i$$
这个式子(对事件样本)的含义是:如果模型认为这个样本出事概率占风险集的 40%(预测),但实际它的 $\delta_i = 1$(事件确实发生了),说明模型给低了,梯度为正,会往提高 $f(x_i)$ 的方向更新。反过来,如果一个删失样本 $\delta_i = 0$ 但被预测了很高的风险,梯度会让它的分数下降。所以这个损失在隐式地做"排序学习",不是在做分类,也不是在做回归。
3. 从传统统计到深度学习:Cox loss 的两种落地姿势
Cox 比例风险模型最早是统计学家 David Cox 在 1972 年提出的,当时求解靠 Newton-Raphson 迭代。现在深度学习流行起来之后,Cox loss 被直接接进了神经网络——把 $f(x)$ 从简单的线性组合 $\beta^T x$ 换成多层感知机的输出,就得到了深度生存模型。但这中间的配方差异值得展开说说,因为你用传统工具包习惯了,转到深度框架时很容易踩坑。
3.1 传统统计方式:R 的 survival 包
在 R 里拟合 Cox 模型,核心代码就那么几行:
r复制library(survival)
model <- coxph(Surv(time, status) ~ age + sex + treatment, data = df)
summary(model)
这个 fit 过程本质上就是在最大化偏似然。survival 包内部处理了很多细节:数据按时间排序、风险集动态构建、打结时间(tied event times)用 Efron 近似或 Breslow 近似修正。你用传统方式时,不需要自己写损失,但代价是你只能拟合线性 $f(x)$,没法捕捉特征之间的复杂非线性交互。
3.2 深度方式:自己写 Cox loss
在 PyTorch 里实现 Cox loss 的经典版本,代码量不大,但细节很多。我调研和复现过多种写法,核心逻辑围绕一个函数:对每个事件样本,计算它在风险集里的 log-sum-exp 值,再取差。
python复制import torch
import torch.nn as nn
class CoxLoss(nn.Module):
def __init__(self):
super().__init__()
def forward(self, log_hazard: torch.Tensor, time: torch.Tensor, event: torch.Tensor) -> torch.Tensor:
"""
log_hazard: 模型输出,shape (batch, 1) 或 (batch,)
time: 观测时间,shape (batch,)
event: 是否发生事件,1=事件,0=删失,shape (batch,)
"""
# 按时间降序排序,方便构造风险集
order = torch.argsort(time, descending=True)
log_hazard = log_hazard[order].squeeze()
time = time[order]
event = event[order]
# 对事件样本计算偏似然
loss = 0.0
n_events = 0
for i in range(len(time)):
if event[i] == 1:
# 风险集:当前样本及其之后(时间更小)的所有样本
risk_set_log_hazards = log_hazard[i:]
# 分子
log_numerator = log_hazard[i]
# 分母: log-sum-exp over risk set
log_denominator = torch.logsumexp(risk_set_log_hazards, dim=0)
loss += log_denominator - log_numerator
n_events += 1
if n_events == 0:
return torch.tensor(0.0, requires_grad=True)
return loss / n_events
这个写法是最直白的教学版,能跑,但有两个问题:一是 Python 循环效率低,二是按事件逐个计算 logsumexp 有大量重复计算。真实项目里建议用向量化版本,或者直接用现有库。
实际上,PyTorch 官方论坛和很多开源库里都有更高效的实现。核心优化思路是:用累积和(cumsum)代替循环。因为风险集是一个"后缀集合"——按时间降序排列后,第 $i$ 个样本的风险集就是 $[i, N)$ 区间,所以分母只需要一次反向累积 logsumexp 就能全部算出来。不过 logsumexp 没有直接的 cumsum 版本,通常做法是先算 logcumsumexp(pytorch 1.8 以上支持 torch.logcumsumexp),然后索引取出。
python复制def cox_loss_vectorized(log_hazard: torch.Tensor, time: torch.Tensor, event: torch.Tensor) -> torch.Tensor:
# 按时间降序排序
order = torch.argsort(time, descending=True)
log_hazard_sorted = log_hazard[order].squeeze()
event_sorted = event[order]
# 反向 logcumsumexp,得到每个位置的风险集 log-sum-exp
# 注意:logcumsumexp 是从前往后累积,这里需要反转
log_hazard_rev = torch.flip(log_hazard_sorted, dims=[0])
log_cumsum_rev = torch.logcumsumexp(log_hazard_rev, dim=0)
log_denominator = torch.flip(log_cumsum_rev, dims=[0])
# 只对事件样本计算损失
event_mask = event_sorted.bool()
loss = (log_denominator[event_mask] - log_hazard_sorted[event_mask]).mean()
return loss
3.3 两种方式怎么选
这条线整理下来,选择依据很清晰:如果特征量不大、可解释性优先、并且你有信心满足比例风险假设,用传统的 coxph 就够了,它自带方差估计、检验、残差诊断,统计严谨性远超自写版本;如果你面对的是高维特征(基因表达、Embedding、图像特征),或者想看看非线性交互能不能提升性能,那就把 Cox loss 接到神经网络里,走深度路线。
我自己的项目经验是:先跑一个传统 Cox 做基线,拿到线性分数和 C-index 底数。然后再上深度 Cox,如果提升不超过 2~3 个点,说明数据本身的信号基本是线性的,深度模型带来的复杂度和过拟合风险不值得。
4. 实践中绕不开的四个细节:打结、删失、评分与批量构造
纯看公式很容易,一到真实数据就全是坑。我在不同数据集上反复跑 Cox loss,总结出几个高频雷区,单独拿出来讲。
4.1 同一时间多个事件:打结(Ties)处理
生存数据里经常出现同一时间多个事件的情况,比如按天记录时,同一天有 3 个患者去世。这时候从数学上严格讲,事件不是逐个发生的,偏似然的定义就变得模糊了。传统统计里有三种近似:
- Breslow 近似:分母重复使用,对所有同时发生的事件,分母都用同一个风险集。计算最快,但事件多时会有偏。
- Efron 近似:处理更精细,分母分多次出现,每次剔除一个已发生事件的样本。精度更高,是 survival 包的默认选项,也是主流推荐。
- 精确法(exact):枚举所有可能的打结顺序,只适合打结数量很少的情况,计算量爆炸。
在深度 Cox loss 的实现里,绝大多数开源库默认采用 Breslow 近似,因为实现简单。但如果你用 PyTorch 自己写,又没处理打结,那你实际上自动采用了 Breslow 近似(因为事件样本逐个循环,分母不动)。如果你的数据里打结比例很高,建议参考 lifelines 或 DeepSurv 的实现,改用 Efron 近似。我自己的经验是:按天记录的数据打结率往往高达 20%-30%,完全忽略会低估标准误、高估 C-index。
4.2 删失样本不是"没用":它们撑起了风险集
很多新手会把删失样本直接丢掉,这是最大的误解。回到公式,删失样本确实不贡献分子项(不参与求和),但它们实打实地出现在分母里。如果没有删失样本,分母的风险集就会少掉一大批"还活着"的人,模型的排序学习就失去了参照系——它就不知道"在某人出事那一刻,还有谁也在风险中"。
举个例子,一个用户在观察期最后一天流失了,你能说他"在第 365 天流失"吗?不能,他只是"在第 365 天还没流失,之后未知"。这个样本的信息是:他在前 364 天都活着。这些信息必须通过风险集进入分母,才能让模型学到"这个人活了这么久,说明他的风险一直不高"。
4.3 C-index 才是真正的评估指标
Cox loss 在训练时最小化负对数偏似然,但业务上大家更关心 C-index(Concordance Index),也就是模型给的风险评分排序和实际事件时间排序的一致程度。C-index 的意义是:随机抽取两个样本,模型给风险更高那个,是否更早发生事件。0.5 等于随机猜,1.0 是完美排序。
C-index 和 Cox loss 的关系是:优化 Cox loss 通常能提高 C-index,但不完全等价。因为 C-index 只看两两排序对,是一个非光滑的 rank-based 指标;而 Cox loss 是一个光滑的替代目标,近似于"加权版的排序损失"。训练结束后,我一般会同时报告 Cox loss 和 C-index,用后者向业务方解释模型效果,前者用于监控训练收敛。
4.4 大批量训练时的风险集构造问题
深度学习训练通常用 mini-batch 随机采样,但对于 Cox loss,这有一个隐患:每个 batch 里随机抽 128 个样本,风险集只包含这 128 个样本,而不是全部训练样本。这会导致分母估计偏差很大,尤其是事件稀疏时,一个 batch 里可能只有 5-6 个事件样本,偏似然的信息量极低。
解决方案大致有三种,按项目需求取舍:
- 全批量训练(full-batch):每次迭代用全部样本构造风险集。数据集小于 5 万时完全可行,梯度稳定,收敛好。
- 按时间分层采样:先把样本按时间区间分组,每个 batch 内保证包含各个时间段的样本,尽量让风险集代表整体。
- 使用 DeepSurv 的采样策略:它默认用全数据集计算风险集,batch 只影响模型前向计算。如果数据太大,可以降级为"近似风险集"——从全量数据中随机抽一部分作为风险集,而不是只从 batch 里抽。
我个人的建议是:如果你的全量数据在 10 万以内,GPU 显存够,直接用 full-batch 实现。我自己在 8 万样本的数据集上用 full-batch 训练,batch size 设成跟全量一样大,每个 epoch 只更新一次梯度,反而比小 batch 更快收敛,损失曲线也平滑得多。
5. 一个完整的 PyTorch 实现:带打结处理与 C-index 评估
从工程角度,把 Cox loss 落成一个可用的训练流程,我习惯把数据预处理、模型、损失、评估串成一个完整模板。这里给出一个我实测过可用的版本,方便直接改来用。
5.1 数据预处理
生存分析的数据格式很固定:每个样本有 time(观测时长)和 event(是否发生事件)。在进入模型之前,需要把特征标准化,时间列和事件列单独拆出来。
python复制import pandas as pd
import numpy as np
from sklearn.preprocessing import StandardScaler
def prepare_survival_data(df, feature_cols, time_col, event_col):
X = df[feature_cols].values
y_time = df[time_col].values.astype(np.float32)
y_event = df[event_col].values.astype(np.float32)
scaler = StandardScaler()
X = scaler.fit_transform(X)
# 按时间降序排序,方便后续损失计算
order = np.argsort(y_time, kind='mergesort')[::-1]
X = X[order]
y_time = y_time[order]
y_event = y_event[order]
return X, y_time, y_event, scaler
注意这里用 kind='mergesort' 保持稳定排序,保证相同时间的样本之间相对顺序不会被打乱。这个细节在打结处理时很重要。
5.2 模型定义
模型就是一个普通的 MLP,输出一个标量 log-risk:
python复制import torch.nn as nn
import torch.nn.functional as F
class SurvivalMLP(nn.Module):
def __init__(self, n_features, hidden_dims=(64, 32), dropout=0.2):
super().__init__()
layers = []
in_dim = n_features
for h_dim in hidden_dims:
layers.append(nn.Linear(in_dim, h_dim))
layers.append(nn.BatchNorm1d(h_dim))
layers.append(nn.ReLU())
layers.append(nn.Dropout(dropout))
in_dim = h_dim
layers.append(nn.Linear(in_dim, 1))
self.net = nn.Sequential(*layers)
def forward(self, x):
return self.net(x)
5.3 带打结处理的 Cox loss(Efron 近似)
为了实现 Efron 近似,需要把同一时间的事件样本分组处理。完整实现可以参考 lifelines 的 C++ 核心逻辑,但 PyTorch 版本我用过一个比较简洁的写法:
python复制def cox_loss_efron(log_hazard, time, event):
# 排序(降序)
order = torch.argsort(time, descending=True, stable=True)
log_hazard = log_hazard[order].squeeze()
time = time[order]
event = event[order]
loss = torch.tensor(0.0, device=log_hazard.device)
n_events = 0
i = 0
N = len(time)
while i < N:
if event[i] == 0:
i += 1
continue
# 找到同一时间的全部样本
t = time[i]
j = i
while j < N and time[j] == t:
j += 1
# 当前时间的事件样本(i 到 j 之间)
event_indices = [k for k in range(i, j) if event[k] == 1]
tie_event = event_indices
m = len(tie_event) # 打结事件数
# 风险集:从 i 到结尾
risk_set = log_hazard[i:]
risk_set_exp = torch.exp(risk_set)
# Efron 近似:对每个第 l 个打结事件,分母剔除前 l 个事件样本的 exp
for l, idx in enumerate(tie_event):
log_numerator = log_hazard[idx]
# 剔除前 l 个事件样本的贡献
removed_exps = 0.0
for k in range(l):
removed_exps += torch.exp(log_hazard[tie_event[k]])
denominator = torch.sum(risk_set_exp) - (l / m) * removed_exps
loss += torch.log(denominator) - log_numerator
n_events += 1
i = j
return loss / n_events if n_events > 0 else torch.tensor(0.0, requires_grad=True)
这个实现是教学向的,效率一般,但逻辑清晰。生产环境建议直接考虑 PyTorch 官方扩展或者优化后的 cumsum 版本。
5.4 C-index 实现
C-index 的计算逻辑不复杂,但要写得高效需要一点技巧:
python复制def concordance_index(time, event, score):
order = np.argsort(time)
time = time[order]
event = event[order]
score = -score[order] # 分数越高,风险越大;但 C-index 习惯用"越高越好"的预测值,取负
n = len(time)
concordant = 0
comparable = 0
for i in range(n):
for j in range(i+1, n):
if event[i] == 1:
comparable += 1
if score[i] > score[j]:
concordant += 1
elif score[i] == score[j]:
concordant += 0.5
return concordant / comparable if comparable > 0 else 0.0
这个双重循环版本数据量一上 5000 就非常慢,实际使用时建议用 scikit-learn 风格的向量化实现,或者直接调 lifelines.utils.concordance_index。写这段只是为了展示思想:C-index 本质上是在"一个发生了事件、一个还没发生(或更晚发生)"的 pair 里,衡量模型的排序是否和真实时间一致。
5.5 训练循环骨架
python复制model = SurvivalMLP(n_features=X.shape[1])
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3)
X_t = torch.tensor(X, dtype=torch.float32)
t_t = torch.tensor(y_time, dtype=torch.float32)
e_t = torch.tensor(y_event, dtype=torch.float32)
for epoch in range(300):
model.train()
log_hazard = model(X_t)
loss = cox_loss_efron(log_hazard, t_t, e_t)
optimizer.zero_grad()
loss.backward()
optimizer.step()
if (epoch + 1) % 50 == 0:
model.eval()
with torch.no_grad():
score = model(X_t).numpy().squeeze()
c_index = concordance_index(y_time, y_event, score)
print(f"Epoch {epoch+1}, Loss: {loss.item():.4f}, C-index: {c_index:.4f}")
6. 训练 Cox loss 时最容易犯的五个错误
这块内容是我在实际项目中踩过坑之后总结的,每个都对应真实事故。列出来可以少走很多弯路。
6.1 特征没有标准化,模型不收敛
Cox loss 的性质和逻辑回归很像,它对特征尺度极其敏感。如果某个特征取值在 0-1 之间,另一个特征值在 0-100000 之间,线性层输出的 log-risk 会被大尺度特征主导,梯度更新时小尺度特征几乎学不到东西。用神经网络时,BatchNorm1d 能缓解一部分,但最稳妥的做法还是进入模型前对特征做标准化。
6.2 事件比例过低,损失被稀释
如果数据集里发生事件的样本只占 5%,那么求和项里只有 5% 的样本在贡献梯度,剩下 95% 的删失样本只通过分母影响梯度。此时模型容易陷入"所有人都低风险"的局部最优。我遇到这种数据时,会在损失函数里对事件样本做权重放大,或者采用case-cohort 采样——确保每个 batch 里事件样本占比不低于 20%。当然,这会引入采样偏差,评估时需要加权校正。
6.3 时间单位不统一,导致风险集错位
比如有的样本时间用天,有的用小时,合并到同一列时不转换单位,排序就乱了。这是个很低级但很常见的错误。我的习惯是,在预处理阶段统一转换到同一个时间单位,并且把时间列和事件列合起来画一下 KM 曲线(Kaplan-Meier)做 sanity check,如果曲线形状看起来不合理,数据预处理大概率有问题。
6.4 把验证集的 C-index 当模型好坏的全部
C-index 只衡量排序一致性,不衡量校准度(calibration)。一个模型可能 C-index 不错,但预测的风险分数绝对值完全偏移。Cox loss 训练出的 $f(x)$ 是相对风险的对数,本身不含基线风险,所以不能直接用 $\exp(f(x))$ 去预测绝对风险。想要绝对风险,需要额外估计累积基线风险 $H_0(t)$,常用方法是 Breslow 估计器。如果你的业务需要输出"未来 30 天流失概率",C-index 再高也不能直接交差,必须校准。
6.5 梯度爆炸:exp 运算的数值稳定性
当模型输出很大(比如超过 15),$\exp(f(x))$ 会溢出到 Inf,反向传播直接 NaN。缓解手段有三个:
- 输出层加个缩放或约束,让 log-risk 控制在合理范围;
- 使用
logsumexp稳定计算分母,而不是先算 exp 再算 log; - 必要时对模型输出做 clip,比如限制在 [-10, 10]。
我早期用不带 logsumexp 的实现,训练到第 40 个 epoch 时 loss 变成 NaN,排查了半天才发现是某几个离群样本的输出值太大,导致风险集里 exp 爆炸。换成 logsumexp 之后稳如老狗。
7. 扩展场景:Cox loss 在竞争风险、时变特征与大规模数据下怎么用
基础版 Cox loss 能覆盖大部分场景,但真实业务里几乎总会遇到几个变种需求。简单聊一下常见的三个方向,方便你遇到时知道该往哪个方向查资料。
7.1 竞争风险(Competing Risks)
如果"事件"不止一种,比如用户流失这个事件可能被"账号注销"打断,或者患者死亡可能与"因其他疾病住院"竞争,标准的 Cox loss 就不够用了。此时需要用 cause-specific Cox 模型 或 Fine-Gray 模型。前者为每种事件单独拟合一个 Cox 模型,把其他类型事件当作删失处理;后者直接建模累积发生函数(CIF)。在 loss 层面,Fine-Gray 的风险集会做加权调整,形式上比标准 Cox loss 复杂一些。
7.2 时变特征(Time-varying Covariates)
标准 Cox 模型假设特征不随时间变化。如果你的特征本身是时变的(比如用户的月消费金额、App 每周活跃天数),最简单的做法是把数据改造成 start-stop 格式:每个样本拆成多行,每行对应一个时间区间,区间内特征不变。此时风险集的定义也要跟着调整,因为一个个体可能出现在多个区间。在深度模型里,更现代的做法是用 Cox-Time 模型,让 $f(x, t)$ 同时依赖时间和特征,属于全参数化方法,能打破比例风险假设。
7.3 百万级样本的深度 Cox
当数据量超过百万,full-batch 训练不再现实。这时通常采用三步走:先用随机采样构造一个 20 万左右的子集,训练一个基础模型;然后通过 在线风险集采样 继续精调——每个 batch 除了自己的样本外,额外随机抽 500-1000 个"参照样本"加入分母风险集;最后用全量数据评估。这种方法在工程上近似了全量风险集,实际效果损失很小,但训练速度提升了几个数量级。
我自己在一个千万级用户数据集上试过,参照样本数从 100 提高到 1000,C-index 提升了 1.2 个点;再继续增大参照样本数,收益就不明显了,但是显存压力大了不少。所以这个超参调到 1000 左右基本够用。
Cox loss 的公式看着简单,里面对应的建模思想和工程细节其实是压缩过的。从偏似然到打结处理到数值稳定实现,每层都有值得展开的内容。把这个公式弄透,生存分析里最核心的一块地基就算是打牢了,后续学 DeepSurv、DeepHit、Cox-Time 这些进阶模型都会顺畅很多——它们本质上都是在这个 base loss 之上做魔改。
