@@ -19,6 +19,7 @@ limitations under the License.
1919#include < torch/nn/functional/pooling.h>
2020
2121#include < algorithm>
22+ #include < cmath>
2223
2324namespace xllm {
2425namespace {
@@ -61,6 +62,12 @@ void RegionECache::init(const DiTCacheConfig& cfg) {
6162 regione_local_edited_global_ids_ = torch::Tensor ();
6263 regione_local_edited_cache_ids_ = torch::Tensor ();
6364 regione_local_image_global_ids_ = torch::Tensor ();
65+ regione_avd_accumulate_ = 1.0 ;
66+ regione_avd_ratio_ = 1.0 ;
67+ regione_avd_raw_gamma_ = 1.0 ;
68+ regione_avd_gamma_ = 1.0 ;
69+ regione_avd_gamma_exponent_ = 1.0 ;
70+ regione_avd_error_ = 0.0 ;
6471 regione_clear_all_prefetch_slots ();
6572}
6673
@@ -96,6 +103,10 @@ bool RegionECache::regione_should_compute_velocity(int64_t step,
96103 double prev_timestep) {
97104 if (!regione_enabled_) {
98105 regione_avd_ratio_ = 1.0 ;
106+ regione_avd_raw_gamma_ = 1.0 ;
107+ regione_avd_gamma_ = 1.0 ;
108+ regione_avd_gamma_exponent_ = 1.0 ;
109+ regione_avd_error_ = 0.0 ;
99110 return true ;
100111 }
101112 // STS / SMS / forced refresh: always run DiT and reset AVD accumulator.
@@ -106,12 +117,17 @@ bool RegionECache::regione_should_compute_velocity(int64_t step,
106117 step <= config_.regione .warmup_steps ) {
107118 regione_avd_accumulate_ = 1.0 ;
108119 regione_avd_ratio_ = 1.0 ;
120+ regione_avd_raw_gamma_ = 1.0 ;
121+ regione_avd_gamma_ = 1.0 ;
122+ regione_avd_gamma_exponent_ = 1.0 ;
123+ regione_avd_error_ = 0.0 ;
109124 return true ;
110125 }
111126
112127 // Diffusers 28-step RegionE transition gamma (inplace.py), 27 values for
113- // transitions between 28 steps. Linearly upsample/downsample onto the
114- // actual (infer_steps - 1) transitions.
128+ // transitions between 28 steps. Linearly sample it onto the actual
129+ // (infer_steps - 1) transitions, then temper the per-step gamma when the
130+ // current run uses more transitions than the reference schedule.
115131 static constexpr double kRegionEGammaRef [] = {
116132 1.0186 , 1.0241 , 1.0236 , 1.0205 , 1.0298 , 1.0221 , 1.0248 , 1.0246 , 1.0269 ,
117133 1.0275 , 1.0323 , 1.0311 , 1.0298 , 1.0353 , 1.0343 , 1.0397 , 1.0387 , 1.0393 ,
@@ -139,25 +155,34 @@ bool RegionECache::regione_should_compute_velocity(int64_t step,
139155 const bool compute =
140156 ((step - config_.regione .warmup_steps ) % interval) == 0 ;
141157 regione_avd_ratio_ = 1.0 ;
158+ regione_avd_raw_gamma_ = 1.0 ;
159+ regione_avd_gamma_ = 1.0 ;
160+ regione_avd_gamma_exponent_ = 1.0 ;
161+ regione_avd_error_ = 0.0 ;
142162 if (compute) regione_avd_accumulate_ = 1.0 ;
143163 return compute;
144164 }
145165
146- // AVDCache (paper Eq.7-9 / inplace.py), step-count agnostic via resampled γ:
166+ // AVDCache (paper Eq.7-9 / inplace.py), timestep-delta normalized γ:
147167 // ratio = gamma(step) * (1 + (t - t_prev) / 1000)
148168 // accumulate *= ratio; error = 1 - accumulate
149- // reuse velocity while error <= cache_threshold and ratio < 1
150- const double gamma = sample_gamma (step);
169+ // reuse velocity while error <= cache_threshold
170+ const double raw_gamma = sample_gamma (step);
171+ const double ref_timestep_delta = 1000.0 / static_cast <double >(kGammaRefLen );
172+ const double local_timestep_delta = std::abs (timestep - prev_timestep);
173+ const double raw_gamma_exponent = local_timestep_delta / ref_timestep_delta;
174+ const double gamma_exponent =
175+ std::max (0.25 , std::min (3.0 , raw_gamma_exponent));
176+ const double gamma = std::pow (raw_gamma, gamma_exponent);
151177 const double ratio = gamma * (1.0 + (timestep - prev_timestep) / 1000.0 );
152178 regione_avd_ratio_ = ratio;
153-
154- if (ratio >= 1.0 ) {
155- regione_avd_accumulate_ = 1.0 ;
156- return true ; // recompute DiT
157- }
179+ regione_avd_raw_gamma_ = raw_gamma;
180+ regione_avd_gamma_ = gamma;
181+ regione_avd_gamma_exponent_ = gamma_exponent;
158182
159183 regione_avd_accumulate_ *= ratio;
160- const double error = 1.0 - regione_avd_accumulate_;
184+ const double error = std::abs (1.0 - regione_avd_accumulate_);
185+ regione_avd_error_ = error;
161186 if (error > static_cast <double >(config_.regione .cache_threshold )) {
162187 regione_avd_accumulate_ = 1.0 ;
163188 return true ; // recompute DiT
@@ -231,6 +256,10 @@ void RegionECache::regione_prepare_inference(
231256 regione_velocity_cache_ = torch::Tensor ();
232257 regione_avd_accumulate_ = 1.0 ;
233258 regione_avd_ratio_ = 1.0 ;
259+ regione_avd_raw_gamma_ = 1.0 ;
260+ regione_avd_gamma_ = 1.0 ;
261+ regione_avd_gamma_exponent_ = 1.0 ;
262+ regione_avd_error_ = 0.0 ;
234263 regione_partial_mode_ = false ;
235264 regione_local_edited_global_ids_ = torch::Tensor ();
236265 regione_local_edited_cache_ids_ = torch::Tensor ();
@@ -793,98 +822,4 @@ std::pair<torch::Tensor, torch::Tensor> RegionECache::regione_patch_img_kv(
793822 return {full_key, full_value};
794823}
795824
796- bool RegionECache::regione_profile_enabled () const {
797- return regione_enabled_ && config_.regione .profile ;
798- }
799-
800- void RegionECache::regione_profile_reset_step (int64_t step,
801- bool partial_step,
802- bool full_step,
803- bool velocity_cache,
804- int64_t step_tokens,
805- int64_t full_tokens) {
806- regione_profile_step_ = step;
807- regione_profile_partial_step_ = partial_step;
808- regione_profile_full_step_ = full_step;
809- regione_profile_velocity_cache_ = velocity_cache;
810- regione_profile_step_tokens_ = step_tokens;
811- regione_profile_full_tokens_ = full_tokens;
812- regione_profile_kv_store_count_ = 0 ;
813- regione_profile_prefetch_issue_count_ = 0 ;
814- regione_profile_prefetch_hit_count_ = 0 ;
815- regione_profile_prefetch_miss_count_ = 0 ;
816- regione_profile_fallback_h2d_count_ = 0 ;
817- regione_profile_patch_scatter_count_ = 0 ;
818- regione_profile_kv_store_cpu_ms_ = 0.0 ;
819- regione_profile_prefetch_issue_ms_ = 0.0 ;
820- regione_profile_prefetch_wait_ms_ = 0.0 ;
821- regione_profile_fallback_h2d_ms_ = 0.0 ;
822- regione_profile_patch_scatter_ms_ = 0.0 ;
823- }
824-
825- void RegionECache::regione_profile_log_step (double transformer_ms,
826- double arp_ms,
827- double scheduler_ms,
828- double total_ms) const {
829- if (!regione_profile_enabled ()) return ;
830- LOG (INFO ) << " [RegionEProfile] step=" << regione_profile_step_ << " mode="
831- << (regione_profile_partial_step_
832- ? " partial"
833- : (regione_profile_full_step_ ? " full" : " reuse" ))
834- << " velocity_cache=" << regione_profile_velocity_cache_
835- << " tokens=" << regione_profile_step_tokens_ << " /"
836- << regione_profile_full_tokens_ << " total_ms=" << total_ms
837- << " transformer_ms=" << transformer_ms
838- << " scheduler_ms=" << scheduler_ms << " arp_ms=" << arp_ms
839- << " kv_store_cpu_ms=" << regione_profile_kv_store_cpu_ms_
840- << " kv_store_count=" << regione_profile_kv_store_count_
841- << " kv_prefetch_issue_ms=" << regione_profile_prefetch_issue_ms_
842- << " kv_prefetch_issue_count="
843- << regione_profile_prefetch_issue_count_
844- << " kv_prefetch_wait_ms=" << regione_profile_prefetch_wait_ms_
845- << " kv_prefetch_hit_count=" << regione_profile_prefetch_hit_count_
846- << " kv_prefetch_miss_count="
847- << regione_profile_prefetch_miss_count_
848- << " kv_fallback_h2d_ms=" << regione_profile_fallback_h2d_ms_
849- << " kv_fallback_h2d_count=" << regione_profile_fallback_h2d_count_
850- << " kv_patch_scatter_ms=" << regione_profile_patch_scatter_ms_
851- << " kv_patch_scatter_count="
852- << regione_profile_patch_scatter_count_;
853- }
854-
855- void RegionECache::regione_profile_add_kv_store (double ms) {
856- if (!regione_profile_enabled ()) return ;
857- regione_profile_kv_store_cpu_ms_ += ms;
858- ++regione_profile_kv_store_count_;
859- }
860-
861- void RegionECache::regione_profile_add_prefetch_issue (double ms) {
862- if (!regione_profile_enabled ()) return ;
863- regione_profile_prefetch_issue_ms_ += ms;
864- ++regione_profile_prefetch_issue_count_;
865- }
866-
867- void RegionECache::regione_profile_add_prefetch_hit (double wait_ms) {
868- if (!regione_profile_enabled ()) return ;
869- regione_profile_prefetch_wait_ms_ += wait_ms;
870- ++regione_profile_prefetch_hit_count_;
871- }
872-
873- void RegionECache::regione_profile_add_prefetch_miss () {
874- if (!regione_profile_enabled ()) return ;
875- ++regione_profile_prefetch_miss_count_;
876- }
877-
878- void RegionECache::regione_profile_add_fallback_h2d (double ms) {
879- if (!regione_profile_enabled ()) return ;
880- regione_profile_fallback_h2d_ms_ += ms;
881- ++regione_profile_fallback_h2d_count_;
882- }
883-
884- void RegionECache::regione_profile_add_patch_scatter (double ms) {
885- if (!regione_profile_enabled ()) return ;
886- regione_profile_patch_scatter_ms_ += ms;
887- ++regione_profile_patch_scatter_count_;
888- }
889-
890825} // namespace xllm
0 commit comments