AI model reliability training method and device based on video memory state perception
Through memory state perception and multi-stream asynchronous transmission technology, the problem of low checkpoint copy performance in large-scale AI model training is solved, the checkpoint storage performance is improved and the training time is shortened, which improves the reliability and efficiency of distributed training.
Patent Information
- Application Number
- CN202510885383.3
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-06-30
- Publication Date
- 2025-10-17
AI Technical Summary
During large-scale AI model training, the performance of copying checkpoints from graphics cards to hosts is poor, resulting in excessive checkpoint storage overhead, affecting training efficiency and performance, and insufficient fault tolerance and recovery capabilities of distributed training systems.
Through video memory status perception technology, the video memory usage during training is analyzed, the checkpoint is divided into two parts, and the multi-stream asynchronous transmission mechanism is used to reduce repeated application and release of video memory, realize parallel transmission and asynchronous copying of checkpoints, and improve training efficiency.
This significantly reduces the overhead of checkpoint storage, shortens the end-to-end training time, and improves the reliability and efficiency of the training process.
Smart Images

Figure CN120806062A_ABST
Abstract
Description
TECHNICAL FIELD
[0001] The present application relates to the field of large model distributed training, and particularly relates to an AI model reliability training method and device based on GPU state perception. BACKGROUND
[0002] AI model training requires long-term multi-round iterative learning. In the training process of large-scale AI models, when a certain iteration of training fails, for example, communication connection is interrupted, hardware fails, code bugs, etc., the training will be immediately interrupted, and all training processes will be abnormally exited. At this time, if there is no recoverable means, fault recovery cannot be performed and training needs to be restarted from the beginning, which greatly wastes training resources and training efficiency. And now, distributed cluster training has become the mainstream method, but the expansion of the cluster also brings higher system failure risk, which makes the fault tolerance recovery capability in the training process particularly critical. At present, the distributed training system based on collective communication architecture generally adopts a periodic checkpoint (Checkpoint) saving mechanism to realize the function of breakpoint resuming. A checkpoint is saved every fixed period, and the model parameters are stored in the disk or remote distributed storage system. When a fault occurs, the fault can be recovered by reading the latest saved checkpoint data, achieving the purpose of breakpoint resuming. If the user wants to recover the training to the step at which the fault occurred, Checkpoint needs to be saved every step, and the performance of copying from the GPU to the host (D2H, Device to Host) is not high. This saving strategy will greatly affect the training performance. Therefore, it is necessary to consider the specific network situation to seek a better saving method under the influence of network training performance, Checkpoint saving frequency and device memory space constraints.
[0003] To address this challenge, researchers have proposed various saving optimization methods, the core of which is to perform checkpoint transmission from GPU to host memory and forward-backward parallel execution in reverse, overlapping transmission time to reduce training waiting. However, in the large model scenario, single-card Checkpoint can be as high as ten G, and the transmission bandwidth from GPU to host memory is not high, which greatly limits the overlap. When checkpoint transmission and forward-backward parallel execution are performed in parallel, the parameter update needs to be performed after the transmission is completed, otherwise the remaining checkpoint may be overwritten when the parameter update is performed directly, resulting in inconsistent checkpoint data. Therefore, this design limitation still causes excessive overhead, which greatly affects the saving performance, further affects the total training time of end-to-end, and even affects the performance of normal training. SUMMARY
[0004] The application aims at the problem of low checkpoint D2H performance and high training memory occupation in large model distributed training, and proposes an AI model reliability training method based on memory state perception, which is used to improve the saving performance in large-scale model distributed training, thereby reducing the additional overhead of checkpoints and shortening the end-to-end training time.
[0005] In order to achieve the above purpose, the technical scheme adopted by the application is as follows:
[0006] An AI model reliability training method based on memory state perception, the steps are as follows:
[0007] Step 1: initialize the first_save (representing the first save) identifier as True, before each training operator execution, first determine whether the current training round needs to be saved through the current training step (step) and the set saving frequency. If the last saving step number plus the saving frequency is greater than or equal to the current step number, the saving operation is triggered, if the saving operation is triggered, otherwise skip saving and execute normally.
[0008] Step 2: if the saving operation is triggered, analyze the network structure of the current training AI model, obtain the current D2H communication bandwidth, single training iteration forward time and checkpoint size, etc. through profiling data. According to the checkpoint size, D2H bandwidth and single iteration forward time, calculate the checkpoint size that can be transmitted in parallel with the training forward. And generate a segmentation strategy.
[0009] Step 2.1: first estimate the total checkpoint size (C model ) of the model through the size of the large model parameter, which is generally equal to the product of the model parameter size and the byte size occupied by the data type; perceive the memory state in the current steady-state training in the training, including total memory, current used memory, remaining memory, etc. Set the single iteration forward time as T fb , the D2H communication bandwidth as B d2h , then the checkpoint size C p that can be transmitted in parallel with the single iteration training forward process is T fb / B d2h , that is, the checkpoint with the size of C p can be transmitted in parallel with the forward, without hindering the parameter update time.
[0010] Step 2.2: the remaining memory (M free) can be calculated by subtracting the current used memory from the total memory, combined with training memory state data, communication bandwidth, checkpoint size and other data, according to the splitting algorithm to generate checkpoint splitting strategy, the checkpoint is divided into two parts, the first part is transmitted in parallel with the previous reverse, and the second part performs memory temporary storage operation and another transmission stream for parallel transmission. First, reserve X size of memory to increase fault tolerance, the current available free memory is M useful = M free -X. Compare the total checkpoint size (C model ) with the available free memory size (M useful ) plus the parallel checkpoint size (C p ), if C model >M useful +C p , then the first part checkpoint size is C first =C model -M useful , and the second part checkpoint size is C second =M useful ; otherwise, C first =C p , C second =C model -C p ; the splitting strategy is generated by the above method;
[0011] Step 3: based on the splitting strategy generated by the current training data, the checkpoint is split, according to the splitting strategy, the checkpoint is divided into two parts, first traverse the entire checkpoint from back to front, constantly add parameter size and judge whether the current size is greater than C second , if greater, record the current parameter index and end the traversal. The current parameter index is regarded as the dividing point of the two parts, and the parameter pointer from the first parameter to the dividing point part is stored in the first part (prev_part), and the parameter pointer from the dividing point (not including the dividing point) to the last parameter part is stored in the second part (storage_part).
[0012] Step 4: After the completion of the segmentation, the storage_part part is checkpointed to the spare video memory. In order to prevent repeated application for temporary storage space when saving multiple times. By the first_save identifier to judge whether it is the first time to save, if it is the first time to save, a new tensor (Tensor) is constructed as the destination Tensor, each parameter is in the form of a tensor, the Tensor in the storage_part is used as the original Tensor, and the newly constructed Tensor is used as the destination Tensor to perform the D2D copy of the video memory. Finally, the copied new Tensor is added to the copy_weights list to complete the video memory temporary storage operation of a single Tensor. If it is not the first time to save, the temporary storage space has been constructed in copy_weights, so the Tensor space constructed for the first time can be reused for video memory temporary storage.
[0013] Step 5: After the completion of the temporary storage, the asynchronous copy operation of the prev_part and storage_part parts is executed in parallel through multiple streams. The transmission operations of the two parts are issued to different streams to independently control the completion time and synchronization position of the two checkpoint parts. Two different streams, CopyDataStream and StorageDataStream, are constructed for asynchronous D2H operation, all Tensors in the prev_part part are traversed, and the asynchronous D2H copy is executed based on CopyDataStream using the underlying D2H copy interface, and the copy completion flag is configured as True. The storage_part is asynchronously copied using StorageDataStream. After issuing the asynchronous copy operation, set copy_action to False to indicate that the current round has executed the asynchronous copy operation.
[0014] Step 6: Update the parameters and complete the training. If the current copy completion flag is True, skip the save operation, and determine whether the current is a parameter update operator according to the operator name before the operator executes. If the parameter update operator is encountered for the first time, the flow synchronization operation is executed to wait for the operation in CopyDataStream to be executed before continuing the training. At this time, the D2H copy of the prev_part part is executed in CopyDataStream, and the prev_part part parameters are not temporarily stored in the video memory. If the parameter update is directly performed without synchronization operation, there will be a data inconsistency problem.
[0015] In a second aspect, the application also provides an AI model reliability training device based on video memory state perception, comprising:
[0016] The acquisition module analyzes the network structure of the current trained AI model, estimates the total checkpoint size of the model, and determines whether to save the current training round through the current training step number and the saving frequency.
[0017] The policy generation module generates a segmentation policy and reserves memory with a size of X to increase fault tolerance if the saving operation is triggered.
[0018] The segmentation module divides the checkpoint into two parts according to the segmentation policy, and obtains the dividing point of the two parts of the checkpoint by calculating the parameter size.
[0019] The temporary storage module temporarily stores part of the checkpoint data in the spare memory, and uses the memory multiplexing technology to temporarily store the checkpoint during the temporary storage.
[0020] The transmission module performs an asynchronous copy operation in parallel through multiple streams after the temporary storage is completed, updates the parameters, and completes the training.
[0021] The present application has the following characteristics and beneficial effects:
[0022] To solve the problem of excessive memory occupation and the influence of weight backup copy on normal training during large-scale model distributed training, the present application proposes a training memory state perception technology to obtain the memory occupation during stable training, performs weight segmentation according to the available memory size and model weight size, temporarily stores part of the checkpoint, and uses memory multiplexing technology to reduce the repeated application and release of memory, thereby reducing the pause time during large model training and the overhead introduced by checkpoint saving.
[0023] To solve the problem of excessive saving overhead caused by low Checkpoint D2H performance, the present application performs asynchronous multi-stream D2H transmission of the two segmented weights, overlaps training and D2H transmission to reduce the additional saving overhead and reduce the end-to-end training time.
[0024] Through the above two reliability optimization methods for distributed large model training, the present application significantly improves the end-to-end time during training and significantly improves the saving performance. BRIEF DESCRIPTION OF DRAWINGS
[0025] Figure 1 The figure is the overall architecture diagram of the embodiment of the present application.
[0026] Figure 2 The figure is the specific flowchart of the embodiment of the present application. DETAILED DESCRIPTION
[0027] The application will be described in detail below with specific examples. The following examples will help those skilled in the art to further understand the application, but do not limit the application in any form. It should be noted that the examples in the application and the features in the examples can be combined with each other without conflict.
[0028] An AI model reliability training method based on video memory state perception, as shown in Figure 1 and Figure 2 , comprising the following steps:
[0029] Step 1, initialize the first_save (representing the first save) identifier as False, before each round of training operator execution, first determine whether the current training round needs to be saved by the current training step (step) and the saving frequency. Get the training information from the python side, if the last save step number plus the saving frequency is greater than or equal to the current step number, trigger the saving operation, and update the last save step to the current step, otherwise skip the saving and execute normally.
[0030] Step 2, if the saving operation is triggered, analyze the network structure of the current training AI model, estimate the total checkpoint size, and obtain the current stable training cluster video memory occupation, D2H communication bandwidth, single training iteration forward time and total checkpoint size, etc. through profiling data. According to the checkpoint size, D2H bandwidth and single iteration forward time, calculate the checkpoint size that can be transmitted in parallel with the pre-training reverse, and generate a checkpoint segmentation strategy according to the calculation.
[0031] Step 2.1, estimate the total checkpoint size (C model ) of the model through the model parameter quantity, which is generally equal to the model parameter quantity (N) multiplied by the byte size of the data type (sizeof(dtype)); In the training, the current stable training video memory state is perceived through profiling, including total video memory, current used video memory, remaining video memory, etc. Set the single iteration forward time as T fb , the D2H bandwidth as B d2h , then the checkpoint size C p that can be transmitted in parallel with the pre-training reverse is T fb / B d2h .
[0032] C model =N*sizeof(dtype)
[0033]
[0034] Step 2.2, the remaining video memory (M free) can be calculated by the total memory minus the current used memory, combined with training memory state data, communication bandwidth, checkpoint size and other data, we generate checkpoint segmentation strategy according to the segmentation algorithm, the checkpoint is divided into two parts, the first part is transmitted in parallel with the previous reverse, the second part performs memory staging operation and another transmission stream is started to perform parallel transmission. First, the checkpoint segmentation strategy first reserves memory with size X to increase fault tolerance, prevents the memory from being full, which may cause subsequent training problems, then the current available free memory is M useful = M free -X. Compare the total checkpoint size (C model ) with the available free memory size (M useful ) plus the parallel checkpoint size (C p ), if C model >M useful +C p , then the first part checkpoint size is C first =C model -M useful , the second part size is C second =M useful ; Otherwise, C first =C p , C second =C model -C p ; The segmentation strategy is generated by the above method;
[0035] Step 3, after generating the segmentation strategy, perform checkpoint segmentation operation according to the segmentation strategy and execute memory staging.
[0036] Step 3.1, according to the segmentation strategy, the checkpoint CKPT is segmented into two parts, first traverse the entire checkpoint from back to front, constantly add parameter size and judge whether the current size is greater than C second , if greater, record the current parameter index and end traversal. The current parameter index is regarded as the dividing point of the two parts, traverse the checkpoint again to store the parameter pointer of the first parameter to the dividing point part in the first part prev_part, and store the parameter pointer of the dividing point (not including the dividing point) to the last parameter part in the second part storage_part.
[0037] Step 3.2, After the completion of the cut, the storage_part part is checkpointed to the spare video memory, and the video memory multiplexing technology is used to store the checkpoint. Determine whether the current is the first save through the first_save identifier, if first_save is True, the current is the first save, then build a new tensor as the destination Tensor, each parameter is in the form of a tensor, the Tensor in storage_part as the original Tensor, the newly constructed Tensor as the destination Tensor executes Device to Device copy, and finally the copied new Tensor is added to the copy_weights list to complete the video memory storage operation. If first_save is False, it means that the current is not the first save, then the copy_weights has constructed the storage space, because the weight size does not change in the whole end-to-end training, so the first constructed Tensor space can be reused for video memory storage. At this time, there is no need to build a new tensor, only to get the corresponding Tensor in copy_weights as the destination Tensor for memory to memory D2D copy. D2D will overwrite the data of the corresponding Tensor in copy_weights.
[0038] Step 4, execute the asynchronous D2H transmission operation of prev_part and storage_part part through multi-stream parallel. Because the checkpoint storage locations of prev_part and storage_part parts are different, it is necessary to issue the transmission operations of the two parts to different streams to independently control the completion time and synchronization location of the two checkpoint parts. Build two streams CopyDataStream and StorageDataStream for asynchronous Device to Host (D2H) operation, traverse all Tensors in prev_part part, use the underlying D2H copy interface to execute asynchronous D2H copy based on CopyDataStream, and configure the copy completion identifier as True. storage_part uses StorageDataStream for asynchronous copy operation. After issuing the asynchronous copy operation, set copy_action to False to indicate that the current round has executed the asynchronous copy operation.
[0039] Step 5: Continue executing the corresponding operator. Before executing an operator, determine whether it is a parameter update operator based on the operator name. If this is the first time encountering a parameter update operator, perform stream synchronization and wait for the operations in CopyDataStream to complete before continuing training. This is because CopyDataStream is currently executing a D2H copy of the prev_part portion, and the prev_part parameters are not temporarily stored in the video memory. Directly updating the parameters will cause data inconsistencies.
[0040] Step 6: When training is complete, start an asynchronous thread to execute the checkpoint disk operation and wait for the StorageDataStream flow to complete execution at the beginning to avoid training failures caused by directly executing serialization and disk before completing the D2H operation.
[0041] Secondly, this application also proposes an AI model reliability training device based on video memory state perception, including:
[0042] Get the module, first initialize the first_save (indicates the first save) identifier, analyze the network structure of the currently trained AI model, and estimate the total checkpoint size of the model (C model ), which is generally equal to the number of model parameters multiplied by the size of the data type in bytes; and during training, profiling is used to obtain data such as total video memory, currently used video memory, remaining video memory, reverse time before a single iteration, and D2H communication bandwidth. The current number of training steps and save frequency are used to determine whether the current training round is saved.
[0043] Strategy generation module, the reverse time before a single iteration is T fb , the D2H communication bandwidth is R d2h , then the checkpoint size C that can be transferred in parallel with the reverse process before single-iteration training p T fb / B d2h , that is, there can be at most C p Checkpoints of this size can be transferred in parallel with the forward and backward passes without hindering the parameter update time. free ) can be calculated by subtracting the currently used video memory from the total video memory. Combined with the training memory status data, communication bandwidth, checkpoint size and other data, a checkpoint splitting strategy is generated according to the splitting algorithm. The checkpoint is split into two parts. The first part is transmitted in parallel with the previous and reverse directions. The second part performs the video memory temporary operation and starts another transmission stream for parallel transmission. Before generating the strategy, the XG large series video memory is first reserved to increase fault tolerance to prevent the video memory from being full and causing possible abnormal problems in subsequent training. The current available free video memory is M useful =M free -X. Compare total checkpoint size (Cmodel ) and the free available GPU memory size (M useful ) plus the parallel checkpoint size (C p ), if C model > M useful + C p , the first part checkpoint size is C first = C model - M useful , and the second part checkpoint size is C second = M useful ; otherwise, C first = C p , and C second = C model - C p ; the split strategy is generated by the above method;
[0044] The split module splits the checkpoint based on the split strategy generated by the current training data, and divides the checkpoint into two parts according to the split strategy. First, the entire checkpoint is traversed from back to front, and the parameter size is constantly added and it is judged whether the current size is greater than C second . If it is greater, the current parameter index is recorded and the traversal is ended. At this time, the dividing point of the two parts has been obtained. The checkpoint is traversed again, and the parameter pointers from the first parameter to the dividing point part are stored in the first part prev_part, and the parameter pointers from the dividing point (not including the dividing point) to the last parameter part are stored in the second part storage_part.
[0045] The temporary storage module temporarily stores the storage_part checkpoint in the free GPU memory after the split. In order to prevent repeated application of temporary storage space when saving multiple times, and to prevent memory from being released in time, causing GPU OOM problem, GPU multiplexing technology is used to temporarily store the checkpoint when temporary storage. It speeds up the temporary storage performance and reduces the possibility of GPU OOM. The first_save identifier is used to judge whether it is the first time to save, if it is the first time to save, a new tensor (Tensor) is constructed as the destination Tensor, and each parameter is in the form of Tensor. The Tensor in storage_part is used as the original Tensor, and the newly constructed Tensor is used as the destination Tensor to perform Device to Device (D2D) copy. Finally, the copied new Tensor is added to the copy_weights list to complete the GPU temporary storage operation of a single Tensor. If it is not the first time to save, the temporary storage space has been constructed in copy_weights, so the Tensor space constructed for the first time can be reused for GPU temporary storage, without the need to apply and frequently release GPU space every time, which speeds up the temporary storage performance and reduces the possibility of GPU OOM.
[0046] The transmission module performs the asynchronous copy operation of the prev_part and storage_part parts in parallel through multiple streams. The two part transmission operations are assigned to different streams to independently control the completion time and synchronization position of the two part checkpoints, accelerate the transmission performance, and ensure data consistency. Two different streams, CopyDataStream and StorageDataStream, are constructed for asynchronous D2H operation, and all tensors in the prev_part part are traversed to perform asynchronous D2H copy based on CopyDataStream using the underlying D2H copy interface.
[0047] Based on the above embodiments, the following comparative cases are provided:
[0048] This embodiment is implemented based on the domestic MindSpore training framework and the domestic Ascend computing device, and the original checkpoint saving method in the MindSpore training framework is used as the Baseline for experimental comparison. The Llama2 13B model is trained using four 910B2 computing cards for testing, the checkpoint size on each card is 13G, and the checkpoint saving overhead is simulated and tested in multiple scenarios such as no free memory state, free memory 7G, and free memory 14G. The experimental results show that, compared with the MindSpore method, the end-to-end checkpoint saving overhead is reduced by up to 43%.
[0049] The basic principles, main features and advantages of the present application are shown and described above. It should be understood by those skilled in the art that the present application is not limited by the above embodiments, and the above embodiments and descriptions in the specification are only preferred examples of the present application and are not intended to limit the present application. Without departing from the spirit and scope of the present application, various changes and improvements can be made to the present application, and these changes and improvements all fall within the scope of the claimed present application. The scope of protection of the present application is defined by the appended claims and their equivalents.
Claims
1. A method for training AI model reliability based on memory state perception, characterized in that: The steps include: Step 1: Before each round of training operator execution, determine whether the current training round is saved based on the current training step number and save frequency; Step 2: If the save operation is triggered, analyze the network structure of the currently trained AI model, calculate the checkpoint size that can be transmitted in parallel with the pre-training reverse, and generate a splitting strategy; Step 3: Split the checkpoint based on the split strategy generated by the current training data; Step 4: After the split is completed, some checkpoints are temporarily stored in the free video memory, and the video memory reuse technology is used to temporarily store the checkpoints; Step 5: After the temporary storage is completed, the asynchronous copy operation is performed in parallel through multiple streams, and the parameters are updated to complete the training.
2. The AI model reliability training method based on video memory state perception according to claim 1 is characterized in that: The specific implementation of step 1 is as follows: initialize the first save identifier to True, and before executing each round of training operator, determine whether the current training round is saved based on the current training step number step and the set save frequency; if the last saved step number plus the save frequency is greater than or equal to the current step number, trigger the save operation; if the save operation is triggered, otherwise skip the save and execute normally.
3. The AI model reliability training method based on video memory state perception according to claim 2 is characterized in that: The specific implementation process of step 2 is as follows: Step 2.1: First, estimate the total checkpoint size C of the model by the size of the large model parameters model , C model It is equal to the number of model parameters multiplied by the size of bytes occupied by the data type; during training, the memory status of the current steady-state training is perceived, including the total memory, the currently used memory, and the remaining memory; the reverse time before a single iteration is set to T fb , the D2H communication bandwidth is B d2h , then the checkpoint size C that can be transferred in parallel with the reverse process before single-iteration training p T fb / B d2h , that is, there are at most C p Checkpoints of varying sizes are transmitted in parallel with forward and backward passes without hindering parameter update time; Step 2.2: Remaining video memory M free The checkpoint splitting strategy is generated by subtracting the currently used video memory from the total video memory, combined with the training memory status data, communication bandwidth, and checkpoint size. The checkpoint is split into two parts according to the splitting algorithm. The first part is transmitted in parallel with the previous and reverse directions, and the second part performs the video memory temporary operation and starts another transmission stream for parallel transmission. The checkpoint splitting strategy first reserves a video memory of size X to increase fault tolerance, so the currently available free video memory is M useful =M free -X; compare total checkpoint size C model and the available free video memory size M useful Add parallel checkpoint size C p , if C model >M useful +C p , then the size of the first part of the checkpoint is C first =C model -M useful , the second part checkpoint size is C second =M useful ; Otherwise there is C first =C p , C second =C model -C p .
4. The AI model reliability training method based on video memory state perception according to claim 3 is characterized in that: The specific implementation of step 3 is as follows: based on the splitting strategy generated by the current training data, the checkpoint is split into two parts according to the splitting strategy. First, the entire checkpoint is traversed from the back to the front, the parameter size is continuously added, and the current size is judged to be greater than C. second If it is greater than, record the current parameter index and end the traversal. Treat the current parameter index as the dividing point between the two parts, traverse the checkpoint again and store the parameter pointer from the first parameter to the dividing point in the first part prev_part, and the parameter pointer from the dividing point to the last parameter in the second part storage_part.
5. The AI model reliability training method based on video memory state perception according to claim 4 is characterized in that: The specific implementation of step 4 is as follows: after the split is completed, the storage_part checkpoints are temporarily stored in the free video memory, and the video memory reuse technology is used to temporarily store the checkpoints during temporary storage. The first save identifier is used to determine whether it is the first save. If it is the first save, a new tensor Tensor is constructed as the destination Tensor. Each parameter exists in the form of a tensor. The Tensor in storage_part is used as the original Tensor, and the newly constructed Tensor is used as the destination Tensor to perform a video memory to video memory D2D copy. Finally, the copied new Tensor is added to the temporary weight copy_weights list to complete the video memory temporary storage operation of a single Tensor; if it is not the first save, the temporary storage space has been constructed in copy_weights, and the Tensor space constructed for the first time is reused for video memory temporary storage.
6. The AI model reliability training method based on video memory state perception according to claim 5 is characterized in that: The step 5 is specifically as follows: Step 5.1: After the temporary storage is completed, asynchronous copy operations are performed on the prev_part and storage_part parts in parallel through multiple streams. The two transfer operations are sent to different streams to independently control the completion time and synchronization position of the two checkpoints. Two different streams are constructed, CopyDataStream and StorageDataStream, for asynchronous D2H operations. All Tensors in the prev_part part are traversed, and the underlying D2H copy interface is used to perform asynchronous D2H copy based on CopyDataStream, and the copy completion flag is set to True. Similarly, the storage_part part uses StorageDataStream for asynchronous copy operations. After sending the asynchronous copy operation, set copy_action to False to indicate that the asynchronous copy operation has been performed in the current round. Step 5.2: If the current copy completion flag is True, skip the save operation. Before executing the operator, determine whether it is a parameter update operator based on the operator name. If it is the first time encountering a parameter update operator, perform stream synchronization and wait for the operation in CopyDataStream to complete before continuing training. At this time, the D2H copy of the prev_part part is executed in CopyDataStream, and the parameters of the prev_part part are not temporarily stored in the video memory.
7. An AI model reliability training device based on video memory state perception, used to implement the AI model reliability training method according to any one of claims 1 to 6, characterized in that: Includes the following modules: Get the module, analyze the network structure of the currently trained AI model, estimate the total checkpoint size of the model, and determine whether the current training round is saved based on the current training steps and save frequency; The strategy generation module triggers the save operation, generates a partitioning strategy, and reserves X amount of video memory to increase fault tolerance. The splitting module divides the checkpoints into two parts according to the splitting strategy and obtains the dividing point of the two parts by calculating the parameter size; The temporary storage module temporarily stores part of the checkpoint data in the free video memory and uses the video memory multiplexing technology to temporarily store the checkpoints; After the temporary storage is completed, the transmission module performs asynchronous copy operations in parallel through multiple streams and updates the parameters to complete the training.