Top-level configuration for compression simulation.
| 184 | |
| 185 | @dataclass |
| 186 | class CompSimConfig: |
| 187 | """Top-level configuration for compression simulation.""" |
| 188 | |
| 189 | enabled: bool = False |
| 190 | quantizer: QuantizerConfig = field(default_factory=QuantizerConfig) |
| 191 | entropy: EntropyConfig = field(default_factory=EntropyConfig) |
| 192 | mask: MaskConfig = field(default_factory=MaskConfig) |
| 193 | |
| 194 | @classmethod |
| 195 | def from_trainer_config(cls, cfg: Any) -> "CompSimConfig": |
| 196 | # New approach: directly use compression_sim_cfg if available |
| 197 | if hasattr(cfg, "compression_sim_cfg"): |
| 198 | return cfg.compression_sim_cfg |
| 199 | |
| 200 | # Legacy support for old field names (for backward compatibility with tests) |
| 201 | enabled = bool(getattr(cfg, "compression_sim", False)) |
| 202 | |
| 203 | entropy_cfg = EntropyConfig( |
| 204 | enabled=bool(getattr(cfg, "entropy_model_opt", False)), |
| 205 | model_type=getattr(cfg, "entropy_model_type", "factorized_model"), |
| 206 | steps=dict(getattr(cfg, "entropy_steps", _default_entropy_steps())), |
| 207 | rd_lambda=getattr(cfg, "rd_lambda", 0.01), |
| 208 | ) |
| 209 | entropy_cfg.ensure_all_attributes() |
| 210 | |
| 211 | mask_cfg = MaskConfig( |
| 212 | enabled=bool(getattr(cfg, "shN_ada_mask_opt", False)), |
| 213 | strategy=getattr(cfg, "shN_ada_mask_strategy", "learnable"), |
| 214 | start_step=getattr(cfg, "ada_mask_steps", 10_000), |
| 215 | cap_max=getattr(getattr(cfg, "strategy", None), "cap_max", None), |
| 216 | ) |
| 217 | grad_threshold = getattr(cfg, "shN_ada_mask_grad_threshold", None) |
| 218 | if grad_threshold is not None: |
| 219 | mask_cfg.gradient.grad_threshold = float(grad_threshold) |
| 220 | |
| 221 | quantizer_cfg = QuantizerConfig() |
| 222 | |
| 223 | return cls( |
| 224 | enabled=enabled, |
| 225 | quantizer=quantizer_cfg, |
| 226 | entropy=entropy_cfg, |
| 227 | mask=mask_cfg, |
| 228 | ) |
| 229 | |
| 230 | def to_dict(self) -> Dict[str, Any]: |
| 231 | return { |
| 232 | "enabled": self.enabled, |
| 233 | "quantizer": { |
| 234 | name: cfg.to_dict() for name, cfg in self.quantizer.attributes.items() |
| 235 | }, |
| 236 | "entropy": { |
| 237 | "enabled": self.entropy.enabled, |
| 238 | "model_type": self.entropy.model_type, |
| 239 | "steps": dict(self.entropy.steps), |
| 240 | "factorized_lr": self.entropy.factorized_lr, |
| 241 | "gaussian_lr": self.entropy.gaussian_lr, |
| 242 | "scheduler_gamma": self.entropy.scheduler_gamma, |
| 243 | "rd_lambda": self.entropy.rd_lambda, |
no outgoing calls