Dialysis competition risk survival prediction method based on VAE and GMM
By applying a combined model of VAE and GMM in the data processing of dialysis patients, the problem of traditional methods being difficult to capture multidimensional physiological changes and identify potential subtypes is solved, achieving more accurate prediction of competitive risk event and support for individualized treatment.
Patent Information
- Application Number
- CN202510247267.9
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-03-04
- Publication Date
- 2025-06-27
AI Technical Summary
The prior art is difficult to effectively capture multidimensional physiological changes in dialysis patients, especially when facing censorship, multi-event or competitive risk scenarios, traditional methods are difficult to fully reveal the dynamics and heterogeneity of dialysis, and it is impossible to identify multiple potential subtypes that may exist.
The dialysis competition risk survival prediction method based on VAE and GMM is adopted, and the high-dimensional clinical features are reconstructed into low-dimensional potential representations through pre-training variational autoencoder, and the patient's latent patterns are clustered in combination with Gaussian mixed model, and the variational autoencoder, mixed Gaussian clustering and multi-event competition risk sub-model are combined to optimize the variational autoencoder, mixed Gaussian clustering and multi-event competition risk sub-model.
This method can automatically reveal the potential distribution structure in the data, divide patients into several clinically significant subtypes or trajectory patterns, improve the accuracy of prediction of different events, and help clinicians quickly identify high-risk or specific characteristic patient groups.
Smart Images

Figure CN120221069A_ABST
Abstract
Description
Technical Field
[0001] The present invention belongs to the field of data processing, and particularly relates to a dialysis competing risks survival prediction method based on VAE and GMM. Background Art
[0002] Dialysis is one of the widely used alternative treatment methods for end-stage renal disease (ESRD). Patients can remove metabolic wastes and excess water in the body through dialysis. Dialysis methods include peritoneal dialysis (PD) and hemodialysis (HD). Compared with hemodialysis, peritoneal dialysis is relatively simple to operate and does not require frequent trips to the dialysis center, which helps to improve the quality of life of patients; while hemodialysis requires the use of a dialysis machine and a vascular access, usually carried out in a dialysis center. However, in clinical practice, the responses of different patients to dialysis vary greatly. Some patients can maintain dialysis adequacy and a good living state for a long time, while some patients may experience multiple events such as technical failure, peritonitis, and cardiovascular complications earlier, and even need to switch to other dialysis methods or undergo kidney transplantation. The factors leading to this difference include the patient's underlying diseases, residual renal function, peritoneal permeability difference (for peritoneal dialysis), nutritional and inflammatory status, types of comorbidities, and lifestyle, etc.
[0003] In previous studies, clinicians usually relied on single or a small number of indicators (such as Kt / V, blood urea nitrogen, serum creatinine, etc.) to judge the prognosis of dialysis patients, or used statistical models (such as Cox regression, Kaplan–Meier curve) to estimate the time of event occurrence. However, these methods are not only difficult to effectively capture the multidimensional physiological changes of patients, but also often require additional assumptions or data splitting when facing censoring, multiple events or competing risks scenarios, and it is difficult to comprehensively reveal the dynamics and heterogeneity of dialysis. Especially when stratifying patient groups, traditional methods cannot identify multiple potential subtypes that may exist, thus limiting the in-depth development of individualized treatment and precise management.
[0004] In view of these limitations, how to simultaneously identify the potential subtypes of dialysis patients and accurately predict the occurrence time of their multiple competing risk events has become an important research direction. Summary of the Invention
[0005] The invention proposes a dialysis competing risks survival prediction method based on VAE and GMM to address the above-mentioned challenges, and specifically adopts the following technical solutions:
[0006] A dialysis competing risks survival prediction method based on VAE and GMM, comprising the following steps:
[0007] S1: Preprocess the follow-up data of dialysis patients, and divide the data into a training set, a validation set, and a test set;
[0008] S2: Pre-train a variational autoencoder to reconstruct high-dimensional clinical features into low-dimensional latent representations, providing a prior for the Gaussian mixture model to cluster the latent patterns of patients;
[0009] S3: Train a joint model, add the KL divergence loss of Gaussian mixture clustering, and jointly optimize the variational autoencoder, Gaussian mixture clustering, and multi-event competing risks sub-model;
[0010] S4: Predict the clustering labels through the logarithmic probability density of Gaussian components, and predict the survival probability through the cumulative risk function.
[0011] S5: Through hyperparameter search, train different combinations of hyperparameters on the training set, evaluate their performance on the validation set, select the best combination of hyperparameters, and finally test the deep learning model trained with the best hyperparameters on the test set.
[0012] Furthermore, in step S1, preprocess the follow-up data of dialysis patients, including:
[0013] Data collection: Select the follow-up records of dialysis patients, including patient basic information, dialysis adequacy assessment, nutritional status assessment, primary diseases, etc.;
[0014] Time and event selection: Select the follow-up records with a baseline time of 3 months and record the event types. If no event occurs during the observation period, it is regarded as right censored;
[0015] Data cleaning: Remove patients with too high missing rates and lack of key indicators, fill in some missing data by the mean imputation method, filter and correct outliers, perform label encoding on categorical variables, and standardize continuous features;
[0016] Divide the dataset: Divide the overall dataset into a training set, a validation set, and a test set according to 6:2:2.
[0017] Furthermore, in step S2, pre-train the variational autoencoder. The variational autoencoder composed of a multi-layer fully connected neural network maps the input feature x to the latent feature z. By learning the mean μ and variance log var parameters of the feature x, generate a feature representation z suitable for subsequent Gaussian mixture model clustering. The loss function of the pre-training process is as follows:
[0018] L pretrain =L recon +L pred
[0019] where L pretrain represents the total loss of the pre-training process, L recon is the loss of reconstructing the feature x, and L predis the prediction loss. Further, the prediction loss function is as follows:
[0020] L pred = αL likelihood + βL ranking + γL calibration
[0021] where α, β, γ are weight parameters, and L likelihood is the log-likelihood loss, L ranking is the ranking loss, and L calibration is the calibration loss.
[0022] Further, in step S3, when training the joint model, the total loss function is as follows:
[0023] L total = w1L recon + w2L KLD + w3L pred
[0024] where L total is the total loss, L recon is the reconstruction loss, w1 is the reconstruction loss weight, and L KLD is the KL divergence loss, w2 is the KL divergence loss weight, and L pred is the prediction loss, and w3 is the prediction loss weight.
[0025] Further, in step S4, the log probability density of the Gaussian component is calculated by the following formula:
[0026]
[0027] where z is the latent feature vector, μ k is the mean vector, Σ k is the covariance matrix, D is the latent space dimension, is the variance of the d-th dimension of the k-th component, z d is the d-th dimensional component of the latent vector, and μ k,d is the mean of the d-th dimension of the k-th component.
[0028] Further, in step S4, the posterior clustering probability is calculated by the following formula:
[0029]
[0030] where η k is the posterior probability that the sample belongs to the k-th Gaussian component, π k is the mixing weight of the k-th Gaussian component, z is the input latent feature vector, and μ k is the mean vector of the k-th Gaussian component, and Σ kis the covariance matrix of the k-th Gaussian component, and K is the total number of Gaussian components.
[0031] Further, in step S4, the class with the maximum posterior probability is selected through the following formula for the final grouping result:
[0032]
[0033] Further, in step S4, the survival probability function is calculated by accumulating the risk function through the following formula:
[0034] S e (t|z) = 1 - Λ e (t|z)
[0035] where S e (t|z) is the survival probability of the latent feature vector z for event e from the baseline time to t, and Λ e (t) is the cumulative risk of event e from the baseline time to t.
[0036] Further, in step S5, the AdamW optimizer is adopted, combined with the learning rate scheduler Cosine AnnealingWarm Restarts to optimize the total loss function, and the optimal hyperparameter combination of the model is iteratively determined through the random search method.
[0037] The beneficial effect of the present invention is to provide a dialysis competing risks survival prediction method based on VAE and GMM, which can learn the deep features in high-dimensional clinical data. By embedding a Gaussian mixture prior in the latent space, it can automatically reveal the latent distribution structure in the data, so as to divide patients into several clinically meaningful subtypes or trajectory patterns, which helps clinicians quickly identify high-risk or specific feature patient groups. At the same time, multiple event parallel branches are introduced and modeled for competing risks, better matching the clinical actual scenario and improving the prediction accuracy for different events. BRIEF DESCRIPTION OF THE DRAWINGS
[0038] In order to more clearly illustrate the technical solutions in the embodiments of the present application or the prior art, the following will briefly introduce the drawings required for the description of the embodiments or the prior art. Obviously, the following drawings are only some embodiments of the present application. For those of ordinary skill in the art, other drawings can be obtained based on these drawings without creative efforts.
[0039] Figure 1 is a schematic flowchart of the dialysis competing risks survival prediction method based on VAE and GMM of the present invention;
[0040] Figure 2Schematic diagram of the framework of the dialysis competing risks survival prediction method based on VAE and GMM of the present invention;
[0041] Figure 3 It is the architecture diagram of the prediction model in the pre-trained variational autoencoder provided by the embodiment of the present invention. Detailed implementation manners
[0042] The embodiments of the present application will be described in detail below. The examples of the embodiments are shown in the accompanying drawings, where the same or similar reference numerals denote the same or similar elements or elements with the same or similar functions from beginning to end. The embodiments described below with reference to the accompanying drawings are exemplary and are intended to explain the present application, but should not be construed as limiting the present application.
[0043] The present application discloses a peritoneal dialysis (PD) competing risks survival prediction method based on a variational autoencoder (VAE) and a Gaussian mixture model (GMM), hereinafter simply referred to as the dialysis competing risks survival prediction method based on VAE and GMM. As Figure 1 shown in and, the dialysis competing risks survival prediction method based on VAE and GMM includes:
[0044] S1: Preprocessing the follow-up data of dialysis patients and dividing the training set and the test set.
[0045] The selected data set is the ZJUPD data set, which is collected from the follow-up records of 2,759 dialysis patients, including patients' basic information (age, gender, weight, height, etc.), dialysis adequacy assessment (Kt / V, Ccr, serum albumin, ultrafiltration volume, etc.), nutritional status assessment, primary diseases, etc.
[0046] Time and event selection: For patient p i , select the follow-up record r of this patient at the baseline time t0 = 3 months i and record the event type e of this patient i . If no event occurs during the observation period, it is regarded as right censored, that is, e i =0. In the present application, the event type is defined as: e i =1: Dialysis technology failure; e i =2: Death due to other reasons.
[0047] Data cleaning: Remove patients with too high missing rates and lack of key indicators, fill in some missing data by the mean imputation method, filter and correct outliers and abnormal measurement results, perform label encoding on categorical variables, and perform standardization processing on continuous features.
[0048] Partition the dataset: Divide the overall dataset into a training set, a validation set, and a test set according to 6:2:2 for subsequent model training, hyperparameter tuning, and final performance evaluation.
[0049] S2: Pretrain the variational autoencoder.
[0050] In this application, the input data is mapped to the latent feature z through the encoder, and z is mapped back to the original data space by the decoder, thereby achieving a compact representation and reconstruction of the data. For the latent feature z, the following assumptions are made:
[0051]
[0052] where π k represents the prior probability of the k-th Gaussian component, μ k and σ k are the mean and variance (or standard deviation) respectively, satisfying
[0053]
[0054] The encoder contains two parallel linear transformation layers, a mean μ calculation layer and a log variance log var calculation layer.
[0055] The formula for the mean μ calculation layer is as follows:
[0056] μ(x) = W μ x + b μ
[0057] where, is the input vector, is the weight matrix, is the bias vector.
[0058] The log variance log var calculation layer formula is as follows:
[0059] log var (x) = W var x + b var
[0060] where, is the input vector, is the weight matrix, is the bias vector.
[0061] After the encoder outputs two vectors, the mean μ and the log variance log var , the reparameterization trick is used to sample the latent variable z from the normal distribution:
[0062] σ = exp(0.5 × log var )
[0063] z = μ + σ ⊙ ∈, ∈ ~ N(0, I)
[0064] Among them, z is the finally obtained latent variable, μ is the mean vector of the output, σ is the standard deviation vector of the output, ∈ is the noise sampled from the standard normal distribution, and ⊙ represents element-wise multiplication. This process not only maps the input data to the low-dimensional latent space, but also ensures the continuity and smoothness of the latent space through the variational inference mechanism of the VAE. At the same time, the reparameterization trick enables the entire network to be trained with end-to-end backpropagation.
[0065] The loss function of the pre-training process is as follows:
[0066] L pretrain = L recon + L pred
[0067] Among them, L pretrain represents the total loss of the pre-training process, L recon is the loss of reconstructing the feature x, and L pred is the prediction loss.
[0068] The reconstruction loss guides the variational autoencoder to learn an effective feature representation of the data. The loss of reconstructing the feature x is as follows:
[0069]
[0070] Among them: x i is the original input feature, is the reconstructed feature, and n is the number of samples
[0071] Based on predicting the risk probabilities of different types of events, the prediction model considers the temporal relationship of the event occurrence and ensures the reliability of the prediction through the calibration loss. The designed architecture is as Figure 3 shown. Using the latent input feature vector z, the conditional risk function is calculated through the trained prediction model:
[0072] λ e (t|z) = PredictLayer e (z)
[0073] Among them, λ e (t|z) is the conditional risk of event e occurring at time t given the latent representation z, e is the event type, t is the time point, z is the latent representation of the sample, and PredictLayer(z) is the deep neural network layer of the prediction model.
[0074] Next, calculate the cumulative risk function:
[0075]
[0076] Among them, Λ e (t) is the cumulative risk of event e from the baseline time to t, and s is the time index.
[0077] The predicted loss function is as follows:
[0078] L pred = αL likelihood + βL ranking + γL calibration
[0079] Among them, α, β, γ are weight parameters, and L pred is the predicted loss, L likelihood is the log-likelihood loss, L ranking is the ranking loss, L callibration is the calibration loss.
[0080] The log-likelihood loss function is as follows:
[0081]
[0082] Among them, e is the index of the event type, k is the actually occurred event type, t is the time point when the event occurs, is the indicator function, which is 1 when the event type k is equal to e, and 0 otherwise, and λ e (t) is the risk probability of event e predicted by the model at time t. The log-likelihood loss encourages the model to give high-risk predictions for actually occurred events and low-risk predictions for unoccurred events.
[0083] The role of the ranking loss is to ensure that the risk predictions are consistent with the actual event occurrence order. Its calculation formula is as follows:
[0084]
[0085] Among them, t i , t j are the event occurrence times of two different samples, Λ(t) is the cumulative risk function, σ is the smoothing parameter, is the indicator function, which is 1 when the event time of sample i is earlier than that of sample j, and 0 otherwise.
[0086] The role of the calibration loss is to ensure that the predicted risk matches the actually observed event incidence rate. Its calculation formula is as follows:
[0087]
[0088] Among them, e is the index of the event type, k is the actually occurred event type, t is the time point, and Λ e (t) is the cumulative risk function of event e, indicating the cumulative probability of event occurrence up to time t, is an indicator function, which is 1 when the event type k is equal to e, and 0 otherwise.
[0089] S3: Train the joint model. Add the KL divergence loss of the Gaussian mixture model on the basis of the pre-trained model parameters. Use the mean μ and variance log obtained by pre-training with the variational autoencoder var to provide initialization parameters for the Gaussian mixture model. Guide the VAE to learn a more structured latent representation through the KL divergence loss, model the latent space, and form a meaningful clustering structure. The KL divergence loss function is as follows:
[0090]
[0091] where, L KLD is the KL divergence loss, K is the number of clusters in the Gaussian mixture model, D is the dimension of the latent space, μ d is the d-th component of the mean of the latent vector output by the encoder, is the d-th component of the variance of the latent vector output by the encoder, μ k,d is the mean of the k-th Gaussian component in the d-th dimension, is the variance of the k-th Gaussian component in the d-th dimension, η k is the posterior probability, representing the probability that the sample belongs to the k-th Gaussian component, p(z) is the marginal distribution probability of the latent variable z, is the expectation with respect to the posterior probability η k of, is the expectation of the log marginal probability.
[0092] Jointly optimize the variational autoencoder, Gaussian mixture clustering, and multi-event competing risks sub-model. The total loss function is as follows:
[0093] L total = w1L recon + w2L KLD + w3L pred
[0094] where: L total is the total loss, L recon is the reconstruction loss, w1 is the reconstruction loss weight, L KLD is the KL divergence loss, w2 is the KL divergence loss weight, L pred is the prediction loss, w3 is the prediction loss weight.
[0095] S4: Predict the clustering labels through the log probability density of the Gaussian components, and predict the survival probability through the cumulative risk function. Obtain the latent input feature vector z after model training, and normalize the mixing weights of the trained Gaussian mixture model. The formula is as follows:
[0096]
[0097] Among them, π k is the mixing weight, μ k is the mean vector, and Σ k is the covariance matrix.
[0098] Calculate the logarithmic probability density of the Gaussian component, and the formula is as follows:
[0099]
[0100] Among them, z is the latent feature vector, μ k is the mean vector, Σ k is the covariance matrix, D is the dimension of the latent space, is the variance of the d-th dimension of the k-th component, z d is the d-th dimensional component of the latent vector, and μ k,d is the mean of the d-th dimension of the k-th component.
[0101] Calculate the posterior clustering probability, and the formula is as follows:
[0102]
[0103] Among them, η k is the posterior probability that the sample belongs to the k-th Gaussian component, π k is the mixing weight of the k-th Gaussian component, z is the input latent feature vector, and μ k is the mean vector of the k-th Gaussian component, and Σ k is the covariance matrix of the k-th Gaussian component, and K is the total number of Gaussian components.
[0104] Select the class with the maximum posterior probability through argmax for the final grouping result, and the formula is as follows:
[0105]
[0106] Calculate the survival probability function through the cumulative risk function:
[0107] S e (t|z) = 1 - Λ e (t|z)
[0108] Among them, S e (t|z) is the survival probability of the latent feature vector z for the event e from the baseline time to t, and Λ e (t) is the cumulative risk of the event e from the baseline time to t.
[0109] S5: Use the random search method to perform hyperparameter search. Randomly search and select hyperparameter combinations in the training set for training and verify the performance on the validation set. Select the hyperparameter combination with the best performance and test it on the test set. It can be understood that for the trained competing risk survival prediction model with the potential feature clustering ability, more accurate prediction results can be obtained in terms of performance. Specifically, the AdamW optimizer is adopted, combined with the learning rate scheduler CosineAnnealingWarm Restarts to optimize the total loss function, and the optimal hyperparameter combination of the model is iteratively determined through the random search method.
[0110] The foregoing has shown and described the basic principles, main features and advantages of the present invention. Those skilled in the art should understand that the above embodiments do not limit the present invention in any form. Any technical solutions obtained by using equivalent replacement or equivalent transformation fall within the protection scope of the present invention.
Claims
1. A dialysis competing risk survival prediction method based on VAE and GMM, characterized in that: The following steps are involved: S1: Preprocess the follow-up data of dialysis patients and divide the data into training set, validation set and test set; S2: Pre-trained variational autoencoder to reconstruct high-dimensional clinical features into low-dimensional latent representations, providing priors for the Gaussian mixture model to cluster the patient's latent patterns; S3: Train the joint model, add the KL divergence loss of mixed Gaussian clustering, and jointly optimize the variational autoencoder, mixed Gaussian clustering, and multi-event competing risk sub-model; S4: Cluster labels are predicted by the log probability density of the Gaussian components, and survival probability is predicted by the cumulative hazard function. S5: Through hyperparameter search, different hyperparameter combinations are trained on the training set, and their performance is evaluated on the validation set. The best hyperparameter combination is selected, and finally the deep learning model trained with the best hyperparameters is tested on the test set.
2. The dialysis competing risk survival prediction method based on VAE and GMM according to claim 1, characterized in that: In step S1, the dialysis patient follow-up data is preprocessed, including: Data collection, follow-up records of dialysis patients were selected, including basic patient information, dialysis adequacy assessment, nutritional status assessment, primary disease, etc.; For time and event selection, follow-up records with a baseline time of 3 months were selected and the event type was recorded. If no event occurred during the observation period, it was considered right-censored; Data cleaning: remove patients with too high missing rates and lack of key indicators, fill in some missing data through mean interpolation, filter and correct outliers, label encode categorical variables, and standardize continuous features; Divide the data set into training set, validation set and test set according to the ratio of 6:2:
2.
3. The dialysis competing risk survival prediction method based on VAE and GMM according to claim 1, characterized in that: In step S2, the variational autoencoder is pre-trained to map the input feature x to the potential feature z through a variational autoencoder composed of a multi-layer fully connected neural network, and the mean μ and variance log of feature x are learned. var Parameters are used to generate feature representation z suitable for subsequent Gaussian mixture model clustering. The loss function of the pre-training process is as follows: L pretrain =L recon +L pred Among them, L pretrain represents the total loss of the pre-training process, L recon is the reconstruction loss of feature x, L pred To predict losses.
4. The dialysis competing risk survival prediction method based on VAE and GMM according to claim 3, characterized in that: The prediction loss function is as follows: L pred =αL likelihood +βL ranking +γL calibration Among them, α, β, γ are weight parameters, L likelihood is the log-likelihood loss, L ranking is the ranking loss, L calibration is the calibration loss.
5. The dialysis competing risk survival prediction method based on VAE and GMM according to claim 1, characterized in that: In step S3, the joint model is trained, and the total loss function is as follows: L total =w1L recon +w2L KLD +w3L pred Among them, L total is the total loss, L recon is the reconstruction loss, w1 is the reconstruction loss weight, L KLD is the KL divergence loss, w2 is the KL divergence loss weight, L pred is the prediction loss, and w3 is the prediction loss weight.
6. The dialysis competing risk survival prediction method based on VAE and GMM according to claim 1, characterized in that: In step S4, the logarithmic probability density of the Gaussian component is calculated by the following formula: Among them, z is the potential feature vector, μ k is the mean vector, Σ k is the covariance matrix, D is the latent space dimension, is the variance of the kth component in the dth dimension, z d is the d-th dimension component of the latent vector, μ k,d is the mean of the kth component in the dth dimension.
7. The dialysis competing risk survival prediction method based on VAE and GMM according to claim 1, characterized in that: In step S4, the posterior clustering probability is calculated by the following formula: Among them, η k is the posterior probability that the sample belongs to the kth Gaussian component, π k is the mixture weight of the kth Gaussian component, z is the potential feature vector of the input, μ k is the mean vector of the kth Gaussian component, Σ k is the covariance matrix of the kth Gaussian component, and K is the total number of Gaussian components.
8. The dialysis competing risk survival prediction method based on VAE and GMM according to claim 1, characterized in that: In step S4, the category with the largest posterior probability is selected for the final grouping result by the following formula:
9. The dialysis competing risk survival prediction method based on VAE and GMM according to claim 1, characterized in that: In step S4, the survival probability function is calculated by the cumulative hazard function using the following formula: S e (t|z)=1-Λ e (t|z) Among them, S e (t|z) is the survival probability of the latent feature vector z for event e from baseline time to t, Λ e (t) is the cumulative risk of event e from baseline time to t.
10. The dialysis competing risk survival prediction method based on VAE and GMM according to claim 1, characterized in that: In step S5, the AdamW optimizer is used in combination with the learning rate scheduler CosineAnnealing Warm Restarts to optimize the total loss function, and the optimal hyperparameter combination of the model is iteratively determined through a random search method.