diff --git a/xllm_ops/mega_gdn_decode/op_kernel/mega_gdn_decode_pto_kernel.h b/xllm_ops/mega_gdn_decode/op_kernel/mega_gdn_decode_pto_kernel.h index 4d08ca4..30af790 100644 --- a/xllm_ops/mega_gdn_decode/op_kernel/mega_gdn_decode_pto_kernel.h +++ b/xllm_ops/mega_gdn_decode/op_kernel/mega_gdn_decode_pto_kernel.h @@ -1568,6 +1568,9 @@ AICORE PTO_INLINE void Run( } #endif } + // Publish all output and persistent-state stores before the next decode + // invocation can consume the in-place Conv/SSM caches. + mega_gdn_decode_pto::SyncAllAiv(); #endif } diff --git a/xllm_ops/mega_gdn_mtp_decode/op_kernel/mega_gdn_mtp_decode_pto_kernel.h b/xllm_ops/mega_gdn_mtp_decode/op_kernel/mega_gdn_mtp_decode_pto_kernel.h index 7dc5910..ef0125b 100644 --- a/xllm_ops/mega_gdn_mtp_decode/op_kernel/mega_gdn_mtp_decode_pto_kernel.h +++ b/xllm_ops/mega_gdn_mtp_decode/op_kernel/mega_gdn_mtp_decode_pto_kernel.h @@ -1741,6 +1741,14 @@ AICORE PTO_INLINE void Run( } } #endif + + // Publish output and checkpoint stores before the next MTP invocation can + // consume the updated Conv/SSM state. +#if defined(PTO_NPU_ARCH_A5) + SyncAllMixA5(); +#elif defined(__DAV_VEC__) || defined(__DAV_C220_VEC__) + mega_gdn_decode_pto::SyncAllAiv(); +#endif } } // namespace mega_gdn_mtp_decode_pto