"""Architecture, CUDA and numerical audits for Dendro Omni.""" from __future__ import annotations from typing import Any import torch from torch import nn from .cell import DendroRecurrentCell from .modeling_dendro_omni import DendroForCausalLM def audit_single_source(model: DendroForCausalLM, *, raise_on_error: bool = True) -> dict[str, Any]: result = model.architecture_audit(raise_on_error=False) source_id = id(model.source_layer.source) parameter_ids = {id(parameter) for parameter in model.parameters()} result["all_parameters_are_source"] = parameter_ids == {source_id} result["no_parameterized_subsystems"] = all( sum(parameter.numel() for parameter in module.parameters(recurse=False)) == 0 for module in model.modules() if not hasattr(module, "source") and module is not model.source_layer ) result["passes"] = bool( result["passes"] and result["all_parameters_are_source"] and result["no_parameterized_subsystems"] ) if raise_on_error and not result["passes"]: raise AssertionError(f"Single-source architecture audit failed: {result}") return result def audit_cuda_compatibility(model: DendroForCausalLM) -> dict[str, Any]: unsupported = [] for name, module in model.named_modules(): if isinstance(module, (nn.RNNBase, nn.EmbeddingBag)): unsupported.append(name) result: dict[str, Any] = { "cuda_available": torch.cuda.is_available(), "functional_sdpa_available": hasattr(torch.nn.functional, "scaled_dot_product_attention"), "unsupported_module_names": unsupported, "source_device": str(model.source_layer.source.device), "source_dtype": str(model.source_layer.source.dtype), } if torch.cuda.is_available(): capability = torch.cuda.get_device_capability() result.update( { "cuda_device": torch.cuda.get_device_name(), "compute_capability": f"{capability[0]}.{capability[1]}", "bf16_supported": torch.cuda.is_bf16_supported(), } ) return result def audit_recurrent_identity(model: DendroForCausalLM) -> dict[str, Any]: cells = [module for module in model.modules() if isinstance(module, DendroRecurrentCell)] return { "physical_cell_count": len(cells), "physical_cell_ids": [id(cell) for cell in cells], "virtual_base_depth": model.config.num_hidden_layers, "max_virtual_depth": model.config.max_total_recurrent_steps, "passes": len(cells) == 1, }