import heapq import warnings from functools import lru_cache from typing import Dict, List, Set, Tuple, Any from dataclasses import dataclass, field from constants import ( TIER_ORDER, TIER_BPW, GGUF_OVERHEAD_FACTOR, CLASS_MAX_TIER, CAN_Q3, ALLOW_LOWER_FLOOR, MTP_DEPLOY_TIER, get_tensor_class, get_tensor_type, is_mtp_tensor, ) # Выносим делитель в константу (8 бит * 1024 байт * 1024 кбайт) BITS_IN_MIB = 8 * 1024 * 1024.0 # Предвычисляем множители размеров для каждого тира (ускорение математики) TIER_SIZE_MULTIPLIER = { tier: (bpw / BITS_IN_MIB) * GGUF_OVERHEAD_FACTOR for tier, bpw in TIER_BPW.items() } # K-quants используют блоки по 256 элементов и требуют выравнивания (padding) 3D-тензоров MoE K_QUANTS = {"Q3_K", "Q4_K", "Q5_K", "Q6_K"} MOE_PAD_TYPES = {"ffn_gate_exps", "ffn_up_exps", "ffn_down_exps", "ffn_down"} MSE_BPW = { "IQ1_S": 1.5625, "IQ2_XXS": 2.0625, "IQ2_XS": 2.3125, "IQ2_S": 2.5, "IQ3_XXS": 3.0625, "Q3_K": 3.4375, "IQ3_S": 3.44, "IQ4_XS": 4.25, "IQ4_NL": 4.25, "Q4_K": 4.50, "Q5_K": 5.50, "Q6_K": 6.5625, "Q8_0": 8.50, "F16": 16.0, } @dataclass(order=True, slots=True) class UpgradeItem: """Элемент очереди апгрейдов. Сравнивается только по neg_utility.""" neg_utility: float group_id: int = field(compare=False) next_tier: str = field(compare=False) cost_delta: float = field(compare=False) def _tier_index(tier: str) -> int: if tier not in TIER_ORDER: raise ValueError(f"Unknown tier: {tier}") return TIER_ORDER.index(tier) def _tier_at(idx: int) -> str: if not (0 <= idx < len(TIER_ORDER)): raise IndexError(f"Tier index {idx} out of range") return TIER_ORDER[idx] def _size_mib(tier: str, n_elements: int) -> float: """Размер тензора в MiB с учётом оверхеда GGUF.""" if n_elements <= 0: return 0.0 return n_elements * TIER_SIZE_MULTIPLIER.get(tier, 0.0) @lru_cache(maxsize=64) def _mse_delta(cur_tier: str, next_tier: str) -> float: return (2 ** (-2 * MSE_BPW[cur_tier])) - (2 ** (-2 * MSE_BPW[next_tier])) def _push_upgrade(group_id: int, group_registry: Dict[int, Tuple[List[str], int, int]], assignments: Dict[str, str], tensor_importance: Dict[str, float], upgrade_queue: List[UpgradeItem], importance_table: Dict[str, Any]): # Храним 3 элемента: имена, реальный размер, размер с padding g_names, g_elements, g_elements_padded = group_registry[group_id] rep_name = g_names[0] cur_tier = assignments[rep_name] cur_idx = _tier_index(cur_tier) rep_info = importance_table.get(rep_name, {}) ttype = rep_info["type"] if "type" in rep_info else get_tensor_type(rep_name) cls = get_tensor_class(ttype) max_tier = CLASS_MAX_TIER.get(cls, "Q8_0") if cur_idx >= _tier_index(max_tier) or cur_idx >= len(TIER_ORDER) - 1: return next_tier = _tier_at(cur_idx + 1) # Выбираем размер в зависимости от того, относится ли тир к K-quants cur_size_g = g_elements_padded if cur_tier in K_QUANTS else g_elements next_size_g = g_elements_padded if next_tier in K_QUANTS else g_elements cost_delta = _size_mib(next_tier, next_size_g) - _size_mib(cur_tier, cur_size_g) quality_delta = _mse_delta(cur_tier, next_tier) if quality_delta <= 0: return if cost_delta == 0: utility_per_mb = float('inf') elif cost_delta < 0: return else: total_g_imp = sum(tensor_importance.get(n, 0) for n in g_names) utility_per_mb = (total_g_imp * quality_delta) / cost_delta heapq.heappush(upgrade_queue, UpgradeItem(-utility_per_mb, group_id, next_tier, cost_delta)) def compute_initial_assignments(non_mtp_names: Set[str], mtp_names: Set[str], importance_table: Dict, allow_q3: bool, is_qat: bool = False) -> Dict[str, str]: assignments = {name: MTP_DEPLOY_TIER for name in mtp_names} for name in non_mtp_names: rep_info = importance_table.get(name, {}) has_imatrix = "importance_mean" in rep_info ttype = rep_info["type"] if "type" in rep_info else get_tensor_type(name) cls = get_tensor_class(ttype) if cls in ("norms", "ssm_params"): assignments[name] = "F16" elif cls == "embd": assignments[name] = "Q4_K" if is_qat else "Q5_K" else: base_floor = ("Q4_K" if cls == "attn_proj" else "IQ4_XS") if is_qat else "Q4_K" if allow_q3 and cls in CAN_Q3 and has_imatrix: assignments[name] = ALLOW_LOWER_FLOOR else: assignments[name] = base_floor return assignments def build_groups(tied_groups: List[List[str]], non_mtp_names: Set[str], ne_map: Dict[str, int], padded_ne_map: Dict[str, int]) -> Dict[int, Tuple[List[str], int, int]]: group_registry = {} assigned_tensors = set() for group_idx, group in enumerate(tied_groups): clean_group = [n for n in group if n in non_mtp_names] if clean_group: g_elements = sum(ne_map.get(n, 0) for n in clean_group) g_elements_padded = sum(padded_ne_map.get(n, 0) for n in clean_group) group_registry[group_idx] = (clean_group, g_elements, g_elements_padded) assigned_tensors.update(clean_group) unassigned_tensors = non_mtp_names - assigned_tensors next_group_idx = len(group_registry) for name in unassigned_tensors: group_registry[next_group_idx] = ([name], ne_map.get(name, 0), padded_ne_map.get(name, 0)) next_group_idx += 1 return group_registry def optimal_classify(importance_table: dict, tied_groups: list, model: dict, target_size_mib: float, allow_q3: bool = False) -> Tuple[dict, dict]: if target_size_mib <= 0: raise ValueError("target_size_mib must be positive") features = model.get("features", {}) has_mtp = features.get("has_mtp", False) n_layers = features.get("n_layers", 31) is_qat = features.get("is_qat", False) model_tensors = model.get("tensors", {}) ne_map = {k: v["n_elements"] for k, v in model_tensors.items()} for tname, info in importance_table.items(): if tname not in ne_map: ne_map[tname] = info["n_elements"] # --- Вычисление MoE Padding (Целочисленное выравнивание) --- moe_d_ff = features.get("moe_intermediate_size", 0) padded_ne_map = dict(ne_map) if moe_d_ff > 0 and moe_d_ff % 256 != 0: aligned_d_ff = ((moe_d_ff + 255) // 256) * 256 for name, n_el in ne_map.items(): ttype = importance_table.get(name, {}).get("type", get_tensor_type(name)) if ttype in MOE_PAD_TYPES: padded_ne_map[name] = (n_el // moe_d_ff) * aligned_d_ff # ----------------------------------------------------------- all_names = set(ne_map.keys()) mtp_names = {n for n in all_names if is_mtp_tensor(n, n_layers)} if has_mtp else set() non_mtp_names = all_names - mtp_names tensor_importance = {} for name in non_mtp_names: raw_imp = importance_table.get(name, {}).get("importance_mean", 0.0) tensor_importance[name] = raw_imp assignments = compute_initial_assignments(non_mtp_names, mtp_names, importance_table, allow_q3, is_qat) group_registry = build_groups(tied_groups, non_mtp_names, ne_map, padded_ne_map) mtp_cost = sum(_size_mib(MTP_DEPLOY_TIER, ne_map.get(n, 0)) for n in mtp_names) effective_target = target_size_mib - mtp_cost current_size = sum( _size_mib(assignments[n], padded_ne_map.get(n, ne_map.get(n, 0)) if assignments[n] in K_QUANTS else ne_map.get(n, 0)) for n in non_mtp_names ) if current_size > effective_target: warnings.warn(f"Initial size {current_size:.1f} MiB already exceeds target {effective_target:.1f} MiB", RuntimeWarning) upgrade_queue = [] for g_id in group_registry: _push_upgrade(g_id, group_registry, assignments, tensor_importance, upgrade_queue, importance_table) while upgrade_queue: item = heapq.heappop(upgrade_queue) g_id, next_tier, cost_delta = item.group_id, item.next_tier, item.cost_delta if cost_delta > 0 and current_size + cost_delta > effective_target: continue for n in group_registry[g_id][0]: assignments[n] = next_tier current_size += cost_delta _push_upgrade(g_id, group_registry, assignments, tensor_importance, upgrade_queue, importance_table) return assignments, padded_ne_map def compute_stats(assignments: dict, ne_map: dict = None, padded_ne_map: dict = None) -> dict: """Собирает статистику по тирам с точным учётом K_QUANTS padding.""" stats = {"by_tier_count": {}, "by_tier_mib": {}, "total_mib": 0.0, "tensor_count": 0} for name, tier in assignments.items(): if not isinstance(tier, str): continue stats["tensor_count"] += 1 stats["by_tier_count"][tier] = stats["by_tier_count"].get(tier, 0) + 1 if ne_map: if padded_ne_map and tier in K_QUANTS: elements = padded_ne_map.get(name, ne_map.get(name, 0)) else: elements = ne_map.get(name, 0) size = _size_mib(tier, elements) stats["by_tier_mib"][tier] = stats["by_tier_mib"].get(tier, 0.0) + size stats["total_mib"] += size return stats