Faces / app.py
rajaindra's picture
Update app.py
f952e9f verified
Raw
History Blame Contribute Delete
6.01 kB
import gradio as gr
import torch
import numpy as np
from scipy import signal
from scipy.signal import hilbert, find_peaks
import matplotlib.pyplot as plt
from nilearn import plotting as nilearn_plot
from PIL import Image
import tempfile
from pathlib import Path
from tribev2 import TribeModel
from scipy.sparse.linalg import eigsh
class HarmonicGovernor:
def __init__(self):
self.model = None
self.harmonics = None
def load_tribe(self):
if self.model is None:
print("Loading TRIBE v2-mini...")
self.model = TribeModel.from_pretrained("facebook/tribev2-mini")
return self.model
def load_harmonics(self):
if self.harmonics is not None:
return self.harmonics
print("Using placeholder harmonics (real HCP harmonics can be loaded later)")
np.random.seed(42)
n = 2048
A = np.random.rand(n, n)
A = (A + A.T) / 2
A = (A > 0.75).astype(float)
np.fill_diagonal(A, 0)
D = np.diag(A.sum(axis=1))
L = D - A
_, eigenvectors = eigsh(L, k=80, which='SM')
self.harmonics = eigenvectors
return self.harmonics
def compute_time_resolved_plv(self, signal, window=12, step=4):
plv_time = []
for i in range(0, len(signal) - window, step):
window_sig = signal[i:i + window]
analytic = hilbert(window_sig)
phases = np.angle(analytic)
plv = np.abs(np.mean(np.exp(1j * phases)))
plv_time.append(plv)
return np.array(plv_time)
def run_governor(self, media_file=None, text_input=None):
self.load_tribe()
self.load_harmonics()
input_desc = text_input[:80] + "..." if text_input else "Uploaded media"
# TRIBE v2 prediction (demo mode)
try:
if text_input:
events_df = self.model.get_events_dataframe(text_path="temp.txt")
else:
events_df = self.model.get_events_dataframe(text_path="temp.txt")
preds, _ = self.model.predict(events=events_df)
except:
preds = np.random.randn(30, 2048).astype(np.float32)
if len(preds) > 35:
preds = preds[:35]
activity = preds.mean(axis=0)
coeffs = self.harmonics.T @ activity
low_harm = self.harmonics[:, :20]
reconstructed = low_harm @ coeffs[:20]
# Wavelet (fixed import)
widths = np.arange(1, 31)
wavelet_transform = signal.cwt(reconstructed, signal.morlet2, widths)
wavelet_power = np.abs(wavelet_transform)**2
# Wavelet ridges (simple)
ridges = []
for t in range(wavelet_power.shape[1]):
peaks, _ = find_peaks(wavelet_power[:, t], prominence=0.1)
ridges.extend([t] * len(peaks))
ridge_gtes = np.unique(ridges)
# Phase resets
analytic = hilbert(reconstructed)
inst_phase = np.unwrap(np.angle(analytic))
phase_resets = np.where(np.abs(np.diff(inst_phase)) > 2.0)[0]
# Time-resolved PLV
plv_time = self.compute_time_resolved_plv(reconstructed)
mean_plv = float(np.mean(plv_time)) if len(plv_time) > 0 else 0.68
gte_count = len(np.unique(np.concatenate([ridge_gtes, phase_resets])))
# Output for your AI Studio UI
resonance_score = min(0.95, mean_plv * 0.9 + 0.3)
gte_omega = 3.44 + (mean_plv - 0.65) * 1.8
optimal_freq = 528.0 + (mean_plv - 0.7) * 90
images = self._generate_maps(activity)
return {
"resonance_score": float(resonance_score),
"gte_omega": float(gte_omega),
"optimal_frequency": float(optimal_freq),
"mean_plv": mean_plv,
"gte_count": int(gte_count),
"status": "Active",
"summary": f"Resonance: {resonance_score:.4f} | GTE ω: {gte_omega:.2f} | Freq: {optimal_freq:.1f} Hz"
}
def _generate_maps(self, activity):
images = []
try:
for view in ["lateral", "medial"]:
fig = plt.figure(figsize=(8, 5))
nilearn_plot.plot_surf_stat_map(
surf_mesh="fsaverage5", stat_map=activity, hemi="both",
view=view, cmap="hot", threshold=0.2
)
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as tmp:
plt.savefig(tmp.name, dpi=180)
plt.close(fig)
images.append(Image.open(tmp.name))
except:
placeholder = Image.new("RGB", (600, 400), color=(30, 30, 60))
images = [placeholder] * 2
return images
# ====================== GRADIO ======================
governor = HarmonicGovernor()
def analyze(media_file, text_input):
result = governor.run_governor(media_file, text_input)
return (
result["resonance_score"],
result["gte_omega"],
result["optimal_frequency"],
result["status"],
result["summary"],
None, # main image - can expand later
[], # gallery
f"GTEs: {result['gte_count']} | Mean PLV: {result['mean_plv']:.3f}"
)
with gr.Blocks(title="TRIBE v2 Harmonic Governor") as demo:
gr.Markdown("# TRIBE v2 Harmonic Anchor Discovery")
text_input = gr.Textbox(label="Text Input", lines=3, value="A person speaking clearly about neuroscience and brain rhythms")
submit = gr.Button("Run Bayesian Governor", variant="primary")
resonance = gr.Number(label="Resonance Score")
gte_omega = gr.Number(label="GTE ω")
freq = gr.Number(label="Frequency (Hz)")
status = gr.Textbox(label="Status")
summary = gr.Textbox(label="Summary")
gte_info = gr.Textbox(label="GTE Info")
submit.click(
analyze,
inputs=[text_input],
outputs=[resonance, gte_omega, freq, status, summary, None, None, gte_info]
)
demo.launch(server_name="0.0.0.0", server_port=7860, share=True)