bai-Mind-8 / v1 /basic_model_test.py
eyupipler's picture
Upload 4 files
c39d6d1 verified
Raw History Blame Contribute Delete
3.02 kB
import torch
import numpy as np
import os
from config import ZUCOConfig
from model import build_model
from transformers import AutoTokenizer
def test_model_basic():
print("\n" + "="*80)
print("bai-Mind-8-v1")
print("="*80 + "\n")
model_path = r"model/path"
config = ZUCOConfig()
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
print(f"[INFO] Device: {device}")
print(f"\n[1/4] Creating Model...")
model = build_model(config)
total_params = sum(p.numel() for p in model.parameters())
trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
print(f" Total Params: {total_params:,}")
print(f" Trainable Params: {trainable_params:,}")
print(f"\n[2/4] Model checkpoint loading...")
print(f" Path: {model_path}")
if os.path.exists(model_path):
checkpoint = torch.load(model_path, map_location=device)
model.load_state_dict(checkpoint['model_state_dict'])
print(f" βœ“ Model loaded successfully!")
if 'epoch' in checkpoint:
print(f" Epoch: {checkpoint['epoch']}")
if 'val_loss' in checkpoint:
print(f" Validation Loss: {checkpoint['val_loss']:.4f}")
else:
print(f" βœ— Model path is not found!")
print(f" Using randomize model")
model.to(device)
model.eval()
# Tokenizer
print(f"\n[3/4] Tokenizer loading...")
tokenizer = AutoTokenizer.from_pretrained(config.language_model_name)
print(f" βœ“ Tokenizer is ready")
print(f"\n[4/4] Test...")
batch_size = 2
eeg_seq_len = 50
eeg = torch.randn(batch_size, eeg_seq_len, config.num_eeg_channels).to(device)
eeg_attention_mask = torch.ones(batch_size, eeg_seq_len).to(device)
print(f" Test EEG shape: {eeg.shape}")
# Inference
with torch.no_grad():
generated_ids = model.generate(
eeg=eeg,
eeg_attention_mask=eeg_attention_mask,
max_length=50,
num_beams=2,
early_stopping=True
)
# Decode
predictions = tokenizer.batch_decode(generated_ids, skip_special_tokens=True)
print(f" βœ“ Model is running succesfully!")
print(f"\n Sample predictions:")
for i, pred in enumerate(predictions):
print(f" [{i+1}] {pred}")
print("\n" + "="*80)
print("TEST RESULT: SUCCESS βœ“")
print("="*80 + "\n")
return {
'status': 'success',
'total_params': total_params,
'trainable_params': trainable_params,
'device': str(device),
'sample_predictions': predictions
}
if __name__ == "__main__":
try:
results = test_model_basic()
print("[OK] Model test completed!")
except Exception as e:
print(f"\n[ERROR]: {e}")
import traceback
traceback.print_exc()