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