Phillnet-Mini-Max / audit.py
ayjays132's picture
Complete Phillnet Mini Text-Vision release v1.1.0
1e114b1 verified
Raw
History Blame Contribute Delete
2.6 kB
"""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,
}