#!/usr/bin/env python3
"""Diagnostic script to inspect actual tensor sizes in proto_lm."""

import sys
sys.path.insert(0, 'C:\\MONIKA')

from salience_os_seed.proto_lm.trainer import ProtoLanguageModel, TrainingConfig

config = TrainingConfig()
config.checkpoint_path = "storage/proto_lm/checkpoint.pt"
model = ProtoLanguageModel(config)

print("=== DIAGNOSTIC: Tensor Dimensions ===")
print(f"vocab.size(): {model.vocab.size()}")
print(f"embed.num_embeddings: {model.embed.num_embeddings}")
print(f"embed.weight.shape: {model.embed.weight.shape}")
print(f"output.out_features: {model.output.out_features}")
print(f"output.weight.shape: {model.output.weight.shape}")
print(f"output.bias.shape: {model.output.bias.shape}")
print(f"\ncore type: {type(model.core)}")
print(f"core config d_model: {model.core.config.d_model}")
print(f"core config state_channels: {model.core.config.state_channels}")

# Try to see internal SASS layer dimensions
try:
    first_block = model.core.blocks[0]
    print(f"\nFirst SASS block info:")
    for name, param in first_block.named_parameters():
        print(f"  {name}: {param.shape}")
except Exception as e:
    print(f"Could not inspect SASS blocks: {e}")

print("\n=== Testing forward pass ===")
import torch
test_input = torch.tensor([[0, 1, 2]], device=model.device)
try:
    embedded = model.embed(test_input)
    print(f"Embedded shape: {embedded.shape}")
    hidden, states = model.core(embedded)
    print(f"Hidden shape after SASS: {hidden.shape}")
    logits = model.output(hidden)
    print(f"Logits shape: {logits.shape}")
    print("✓ Forward pass succeeded")
except Exception as e:
    print(f"✗ Forward pass failed: {e}")
    import traceback
    traceback.print_exc()
