Spaces:
Paused
Paused
| # 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 | |