# Copyright (c) 2026 SandAI. All Rights Reserved. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, software # distributed under the License is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. from typing import List import numpy as np from magi_compiler.utils import magi_logger def exponential_aligned_sampler(min_val: int, max_val: int, num_samples: int, align: int = 8) -> List[int]: if min_val >= max_val: raise ValueError(f"最小值({min_val})必须小于最大值({max_val})") if num_samples < 2: raise ValueError(f"采样个数({num_samples})需≥2(至少包含min/max)") if align <= 0: raise ValueError(f"对齐倍数({align})必须为正整数") if align > (max_val - min_val): raise ValueError(f"对齐倍数({align})过大,超过范围跨度({max_val - min_val})") if num_samples > ((max_val - min_val) // align + 1): raise ValueError(f"采样个数({num_samples})过大,无法在范围内生成足够对齐值") aligned_min = ((min_val + align - 1) // align) * align aligned_max = (max_val // align) * align if aligned_min == aligned_max: raise ValueError(f"对齐后min/max均为{aligned_min},请调整align或输入范围") raw_samples = np.logspace(np.log(aligned_min), np.log(aligned_max), num=num_samples, base=np.e) aligned_samples = (np.round(raw_samples / align) * align).astype(int) final_samples = sorted(list(dict.fromkeys(aligned_samples.tolist()))) if len(final_samples) < num_samples: final_samples = np.linspace(aligned_min, aligned_max, num=num_samples) final_samples = (np.round(final_samples / align) * align).astype(int).tolist() magi_logger.info("生成对齐采样点: %s", final_samples, rank=0) return final_samples