Views
No views yet
1import torch
2from transformers import ASTForAudioClassification, ASTFeatureExtractor
3
4# 1. Configuration - Choose which model to load
5repo_id = "MIT/ast-finetuned-audioset-10-10-0.4593" # Use original AST
6# repo_id = "willychenwii/pig-noise-ast-finetuned" # Finetuned model for breathing noise clasiification
7# repo_id = "willychenwii/pig-condition-ast-finetuned" # Finetuned model for breathing condition classification
8
9print(f"--- Loading model from Hub: {repo_id} ---")
10model = ASTForAudioClassification.from_pretrained(repo_id)
11feature_extractor = ASTFeatureExtractor.from_pretrained(repo_id)
12
13# 2. Inspect Model Parameters
14total_params = sum(p.numel() for p in model.parameters())
15trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
16
17print(f"Total Parameters: {total_params:,}")
18print(f"Trainable Parameters: {trainable_params:,}")
19print(f"Hidden Layer Size (d_model): {model.config.hidden_size}")
20
21# 3. Dummy Input (10 seconds of silence at 16kHz)
22sample_rate = 16000
23duration = 10
24dummy_audio = torch.zeros(sample_rate * duration)
25
26# Feature extraction
27inputs = feature_extractor(dummy_audio, sampling_rate=sample_rate, return_tensors="pt")
28print(f"\nInput Shape (to Transformer): {inputs.input_values.shape}")
29# Expected: [batch, time_frames, freq_bins] -> [1, 1024, 128]
30
31# 4. Inference & Hidden States
32model.eval()
33with torch.no_grad():
34 outputs = model(**inputs, output_hidden_states=True)
35
36# 5. Output Shapes
37logits = outputs.logits
38last_hidden_state = outputs.hidden_states[-1]
39
40print(f"Output (Logits) Shape: {logits.shape}") # [1, num_labels]
41print(f"Last Hidden Layer Embedding Size: {last_hidden_state.shape}")
42# Expected: [batch, sequence_length, hidden_size] -> [1, 1214, 768]