Skip to content

Commit 981dd4c

Browse files
committed
refactor: simplify regione avd gamma and remove profile flag.
1 parent 384bfbd commit 981dd4c

8 files changed

Lines changed: 49 additions & 237 deletions

File tree

xllm/core/common/global_flags.h

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -365,7 +365,6 @@ DECLARE_double(dit_regione_region_threshold);
365365
DECLARE_double(dit_regione_cache_threshold);
366366
DECLARE_bool(dit_regione_use_avd_gamma);
367367
DECLARE_bool(dit_regione_erosion_dilation);
368-
DECLARE_bool(dit_regione_profile);
369368

370369
DECLARE_bool(dit_sp_communication_overlap);
371370

xllm/core/framework/config/dit_config.cpp

Lines changed: 2 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -62,10 +62,10 @@ DEFINE_double(dit_regione_region_threshold,
6262
"RegionE: cosine threshold for adaptive region partition.");
6363

6464
DEFINE_double(dit_regione_cache_threshold,
65-
0.03,
65+
0.02,
6666
"RegionE: AVDCache error threshold δ (paper Eq.8). "
6767
"Reuse velocity while 1-accumulate <= threshold. "
68-
"Qwen-Image-Edit default in RegionE inplace.py is 0.03.");
68+
"Default is 0.02.");
6969

7070
DEFINE_bool(dit_regione_use_avd_gamma,
7171
true,
@@ -78,11 +78,6 @@ DEFINE_bool(dit_regione_erosion_dilation,
7878
true,
7979
"RegionE: enable erosion/dilation for region mask cleanup.");
8080

81-
DEFINE_bool(dit_regione_profile,
82-
false,
83-
"RegionE: print per-step timing breakdown for partial/full DiT and "
84-
"K/V CPU offload.");
85-
8681
DEFINE_bool(dit_sp_communication_overlap,
8782
true,
8883
"Communication & Computation overlap for sequence parallel");
@@ -166,7 +161,6 @@ void DiTConfig::from_flags() {
166161
XLLM_CONFIG_ASSIGN_FROM_FLAG(dit_regione_cache_threshold);
167162
XLLM_CONFIG_ASSIGN_FROM_FLAG(dit_regione_use_avd_gamma);
168163
XLLM_CONFIG_ASSIGN_FROM_FLAG(dit_regione_erosion_dilation);
169-
XLLM_CONFIG_ASSIGN_FROM_FLAG(dit_regione_profile);
170164
XLLM_CONFIG_ASSIGN_FROM_FLAG(dit_sp_communication_overlap);
171165
XLLM_CONFIG_ASSIGN_FROM_FLAG(dit_debug_print);
172166
XLLM_CONFIG_ASSIGN_FROM_FLAG(dit_laser_attention_enabled);
@@ -198,7 +192,6 @@ void DiTConfig::from_json(const JsonReader& json) {
198192
XLLM_CONFIG_ASSIGN_FROM_JSON(dit_regione_cache_threshold);
199193
XLLM_CONFIG_ASSIGN_FROM_JSON(dit_regione_use_avd_gamma);
200194
XLLM_CONFIG_ASSIGN_FROM_JSON(dit_regione_erosion_dilation);
201-
XLLM_CONFIG_ASSIGN_FROM_JSON(dit_regione_profile);
202195
XLLM_CONFIG_ASSIGN_FROM_JSON(dit_sp_communication_overlap);
203196
XLLM_CONFIG_ASSIGN_FROM_JSON(dit_debug_print);
204197
XLLM_CONFIG_ASSIGN_FROM_JSON(dit_laser_attention_enabled);
@@ -246,8 +239,6 @@ void DiTConfig::append_config_json(nlohmann::ordered_json& config_json) const {
246239
config_json, default_config, dit_regione_use_avd_gamma);
247240
APPEND_CONFIG_JSON_VALUE_IF_NOT_DEFAULT(
248241
config_json, default_config, dit_regione_erosion_dilation);
249-
APPEND_CONFIG_JSON_VALUE_IF_NOT_DEFAULT(
250-
config_json, default_config, dit_regione_profile);
251242
APPEND_CONFIG_JSON_VALUE_IF_NOT_DEFAULT(
252243
config_json, default_config, dit_sp_communication_overlap);
253244
APPEND_CONFIG_JSON_VALUE_IF_NOT_DEFAULT(

xllm/core/framework/config/dit_config.h

Lines changed: 1 addition & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -56,7 +56,6 @@ class DiTConfig final {
5656
"dit_regione_cache_threshold",
5757
"dit_regione_use_avd_gamma",
5858
"dit_regione_erosion_dilation",
59-
"dit_regione_profile",
6059
"dit_sp_communication_overlap",
6160
"dit_debug_print",
6261
"dit_laser_attention_enabled",
@@ -97,14 +96,12 @@ class DiTConfig final {
9796

9897
PROPERTY(double, dit_regione_region_threshold) = 0.80;
9998

100-
PROPERTY(double, dit_regione_cache_threshold) = 0.03;
99+
PROPERTY(double, dit_regione_cache_threshold) = 0.02;
101100

102101
PROPERTY(bool, dit_regione_use_avd_gamma) = true;
103102

104103
PROPERTY(bool, dit_regione_erosion_dilation) = true;
105104

106-
PROPERTY(bool, dit_regione_profile) = false;
107-
108105
PROPERTY(bool, dit_sp_communication_overlap) = true;
109106

110107
PROPERTY(bool, dit_debug_print) = false;

xllm/core/framework/dit_cache/dit_cache_config.h

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -63,14 +63,12 @@ struct RegionEOptions : public DiTBaseCacheOptions {
6363
int64_t tail_steps = 1;
6464
std::vector<int64_t> refresh_steps = {16};
6565
float region_threshold = 0.80f;
66-
// AVDCache δ in paper Eq.8/9 / inplace.py cache_threshold (Qwen default
67-
// 0.03).
68-
float cache_threshold = 0.03f;
66+
// AVDCache δ in paper Eq.8/9.
67+
float cache_threshold = 0.02f;
6968
// Use fitted γ_t AVDCache (paper) instead of fixed skip_interval.
7069
bool use_avd_gamma = true;
7170
// Enable erosion/dilation morphological cleanup after ARP mask selection.
7271
bool erosion_dilation = true;
73-
bool profile = false;
7472
};
7573

7674
struct ResidualCacheOptions {

xllm/core/framework/dit_cache/regione.cpp

Lines changed: 40 additions & 105 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@ limitations under the License.
1919
#include <torch/nn/functional/pooling.h>
2020

2121
#include <algorithm>
22+
#include <cmath>
2223

2324
namespace xllm {
2425
namespace {
@@ -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

xllm/core/framework/dit_cache/regione.h

Lines changed: 4 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -135,24 +135,6 @@ class RegionECache {
135135
const torch::Tensor& image_rope,
136136
int64_t key_len) const;
137137

138-
bool regione_profile_enabled() const;
139-
void regione_profile_reset_step(int64_t step,
140-
bool partial_step,
141-
bool full_step,
142-
bool velocity_cache,
143-
int64_t step_tokens,
144-
int64_t full_tokens);
145-
void regione_profile_log_step(double transformer_ms,
146-
double arp_ms,
147-
double scheduler_ms,
148-
double total_ms) const;
149-
void regione_profile_add_kv_store(double ms);
150-
void regione_profile_add_prefetch_issue(double ms);
151-
void regione_profile_add_prefetch_hit(double wait_ms);
152-
void regione_profile_add_prefetch_miss();
153-
void regione_profile_add_fallback_h2d(double ms);
154-
void regione_profile_add_patch_scatter(double ms);
155-
156138
private:
157139
void regione_select_regions(const torch::Tensor& sample,
158140
const torch::Tensor& model_output,
@@ -236,23 +218,10 @@ class RegionECache {
236218
torch::Tensor regione_velocity_cache_;
237219
double regione_avd_accumulate_ = 1.0;
238220
double regione_avd_ratio_ = 1.0;
239-
int64_t regione_profile_step_ = -1;
240-
bool regione_profile_partial_step_ = false;
241-
bool regione_profile_full_step_ = false;
242-
bool regione_profile_velocity_cache_ = false;
243-
int64_t regione_profile_step_tokens_ = 0;
244-
int64_t regione_profile_full_tokens_ = 0;
245-
int64_t regione_profile_kv_store_count_ = 0;
246-
int64_t regione_profile_prefetch_issue_count_ = 0;
247-
int64_t regione_profile_prefetch_hit_count_ = 0;
248-
int64_t regione_profile_prefetch_miss_count_ = 0;
249-
int64_t regione_profile_fallback_h2d_count_ = 0;
250-
int64_t regione_profile_patch_scatter_count_ = 0;
251-
double regione_profile_kv_store_cpu_ms_ = 0.0;
252-
double regione_profile_prefetch_issue_ms_ = 0.0;
253-
double regione_profile_prefetch_wait_ms_ = 0.0;
254-
double regione_profile_fallback_h2d_ms_ = 0.0;
255-
double regione_profile_patch_scatter_ms_ = 0.0;
221+
double regione_avd_raw_gamma_ = 1.0;
222+
double regione_avd_gamma_ = 1.0;
223+
double regione_avd_gamma_exponent_ = 1.0;
224+
double regione_avd_error_ = 0.0;
256225
std::vector<torch::Tensor> regione_k_cache_cpu_;
257226
std::vector<torch::Tensor> regione_v_cache_cpu_;
258227
std::vector<torch::Tensor> regione_cond_k_cache_cpu_;

xllm/core/runtime/dit_worker_impl.cpp

Lines changed: 0 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -111,8 +111,6 @@ DiTCacheConfig parse_dit_cache_from_flags() {
111111
::xllm::DiTConfig::get_instance().dit_regione_use_avd_gamma();
112112
cache_config.regione.erosion_dilation =
113113
::xllm::DiTConfig::get_instance().dit_regione_erosion_dilation();
114-
cache_config.regione.profile =
115-
::xllm::DiTConfig::get_instance().dit_regione_profile();
116114
} else if (::xllm::DiTConfig::get_instance().dit_cache_policy() == "None") {
117115
cache_config.selected_policy = PolicyType::None;
118116
}

0 commit comments

Comments
 (0)