duongthienz's picture
Update state.py
3380fea verified
Raw
History Blame
26 kB
"""
state.py — Streamlit session-state management and UI callbacks.
Covers:
- init_session_state()
- Speaker rename helpers (get_display_name, apply_speaker_renames_to_df)
- Category callbacks (addCategory, removeCategory, updateCategoryOptions)
- Global rename callbacks (addGlobalRename, removeGlobalRename, on_grename_change,
apply_inline_rename)
- File-switch callback (updateMultiSelect)
- analyze() — builds and caches all DataFrames for a single file
- convert_df(), printV()
"""
import copy
import traceback
import pandas as pd
import streamlit as st
import sonogram_utility as su
import utils
# ---------------------------------------------------------------------------
# Logging
# ---------------------------------------------------------------------------
verbosity = 4 # 0=None 1=Low 2=Medium 3=High 4=Debug
def printV(message, level):
if verbosity >= level:
print(message)
# ---------------------------------------------------------------------------
# Session state initialisation
# ---------------------------------------------------------------------------
def init_session_state():
"""Idempotently initialise every session-state key the app needs."""
defaults = {
"results": {}, # {filename: (annotations, totalSeconds)}
"speakerRenames": {}, # {filename: {speaker: name}}
"summaries": {}, # {filename: {df2, df3, ...}}
"categories": ["Instructor", "Students"],
"categorySelect": [[], []], # [[token, ...], ...] one list per category, tokens = "fname: SPEAKER_##"; starts with 2 lists for Instructor/Students
"removeCategory": None,
"resetResult": False,
"unusedSpeakers": {}, # {filename: [speaker, ...]}
"file_names": [],
"valid_files": [],
"file_paths": {}, # {filename: path}
"showSummary": "No",
"speakerClips": {}, # {filename: {speaker: wav_bytes}}
"speakerSegments": {}, # {filename: {speaker: [(start,end), ...]}}
"speakerWaveforms": {}, # {filename: (waveform_tensor, sample_rate)}
"globalRenames": [], # [{"name": str, "speakers": ["file: SPEAKER_##", ...]}]
"analyzeAllToggle": False,
}
for key, value in defaults.items():
if key not in st.session_state:
st.session_state[key] = value
# ---------------------------------------------------------------------------
# Display-name helpers
# ---------------------------------------------------------------------------
def get_display_name(speaker, fileName):
"""Return the user-assigned display name for a speaker, or the original label.
Role assignments (categorySelect) are intentionally excluded — roles are for
grouping in charts, not for renaming speakers.
"""
return st.session_state.speakerRenames.get(fileName, {}).get(speaker, speaker)
def apply_speaker_renames_to_df(df, fileName, column="task"):
"""Replace SPEAKER_## labels in a DataFrame column with display names."""
if column not in df.columns:
return df
df = df.copy()
df[column] = df[column].apply(lambda s: get_display_name(s, fileName))
return df
@st.cache_data
def convert_df(df):
return df.to_csv(index=False).encode("utf-8")
def build_all_csv_zip():
"""Build an in-memory ZIP containing one CSV per analyzed file.
Applies the same transformations as the single-file download:
drop Task, rename Resource -> Speaker, sort by Start, add Role.
Returns raw ZIP bytes ready for st.download_button.
"""
import io
import zipfile
buf = io.BytesIO()
with zipfile.ZipFile(buf, "w", zipfile.ZIP_DEFLATED) as zf:
for fname, result in st.session_state.results.items():
if len(result) != 2:
continue
try:
annotation, _ = result
currDF, _ = su.annotationToSimpleDataFrame(annotation)
# Add Role column against raw SPEAKER_## labels BEFORE renames
raw_to_role = {
token.split(": ", 1)[1]: st.session_state.categories[i]
for i, tokens in enumerate(st.session_state.categorySelect)
for token in tokens
if token.startswith(f"{fname}: ")
}
currDF = currDF.copy()
currDF["Role"] = currDF["Resource"].map(raw_to_role).fillna("")
# Apply speaker renames after Role is set
renames = st.session_state.speakerRenames.get(fname, {})
if "Resource" in currDF.columns:
currDF["Resource"] = currDF["Resource"].apply(
lambda s: renames.get(s, s)
)
currDF = currDF.drop(columns=["Task"], errors="ignore")
currDF = currDF.rename(columns={"Resource": "Speaker"})
if "Start" in currDF.columns:
currDF = currDF.sort_values("Start").reset_index(drop=True)
plain_name = fname.rsplit(".", 1)[0]
zf.writestr(
f"sonogram-analysis-{plain_name}.csv",
currDF.to_csv(index=False)
)
except Exception as e:
print(f"build_all_csv_zip: skipping {fname}{e}")
buf.seek(0)
return buf.read()
# ---------------------------------------------------------------------------
# Category callbacks
# ---------------------------------------------------------------------------
def addCategory():
new = st.session_state.categoryInput.strip()
if not new:
return
st.toast(f"Adding {new}")
st.session_state.categories.append(new)
st.session_state.categorySelect.append([])
st.session_state.pop(f"multiselect_{new}", None)
st.session_state.categoryInput = ""
def removeCategory(index):
name = st.session_state.categories[index]
st.toast(f"Removing {name}")
st.session_state.pop(f"multiselect_{name}", None)
del st.session_state.categories[index]
del st.session_state.categorySelect[index]
def updateCategoryOptions(token_display_map=None):
"""Store tokens ('fname: SPEAKER_##') in the global categorySelect lists.
token_display_map: dict {display_label -> raw_token} passed from ui.py.
Widget keys hold display labels; categorySelect must hold raw tokens.
"""
if st.session_state.resetResult:
return
display_to_raw = token_display_map or {}
# Guard: ensure categorySelect has exactly one slot per category.
# Rapid interactions (e.g. addCategory firing mid-callback) can leave
# the two lists temporarily out of sync.
while len(st.session_state.categorySelect) < len(st.session_state.categories):
st.session_state.categorySelect.append([])
while len(st.session_state.categorySelect) > len(st.session_state.categories):
st.session_state.categorySelect.pop()
for i, category in enumerate(st.session_state.categories):
ms_key = f"multiselect_{category}"
display_vals = list(st.session_state.get(ms_key, []))
raw_vals = [display_to_raw.get(t, t) for t in display_vals]
st.session_state.categorySelect[i] = raw_vals
# Recompute unusedSpeakers for all files
all_assigned_tokens = {
token
for tokens in st.session_state.categorySelect
for token in tokens
}
for fname, result in st.session_state.results.items():
if len(result) != 2:
continue
try:
annotation, _ = result
unused = [
sp for sp in annotation.labels()
if f"{fname}: {sp}" not in all_assigned_tokens
]
st.session_state.unusedSpeakers[fname] = unused
except Exception:
pass
# ---------------------------------------------------------------------------
# Global rename callbacks
# ---------------------------------------------------------------------------
# ---------------------------------------------------------------------------
# Global rename helpers
# ---------------------------------------------------------------------------
def _global_rename_key(index):
return f"grename_speakers_{index}"
def _write_rename(token, name):
"""Write name into speakerRenames for a single token (fname: SPEAKER_##).
If name is empty, clears the entry (revert to raw label).
Silently ignores tokens that don't match a known file.
"""
if ": " not in token:
return
fname, raw_sp = token.split(": ", 1)
if fname not in st.session_state.speakerRenames:
return # token references an unknown file — ignore
if name:
st.session_state.speakerRenames[fname][raw_sp] = name
else:
st.session_state.speakerRenames[fname].pop(raw_sp, None)
def addGlobalRename():
new_name = st.session_state.globalRenameInput.strip()
if not new_name:
return
for entry in st.session_state.globalRenames:
if entry["name"] == new_name:
st.toast(f"'{new_name}' already exists in the rename list")
st.session_state.globalRenameInput = ""
return
st.toast(f"Adding rename '{new_name}'")
st.session_state.globalRenames.append({"name": new_name, "speakers": []})
st.session_state.globalRenameInput = ""
def removeGlobalRename(index):
entry = st.session_state.globalRenames[index]
st.toast(f"Removing rename '{entry['name']}'")
# Revert every speaker that belonged to this entry
for token in entry["speakers"]:
_write_rename(token, "")
st.session_state.pop(_global_rename_key(index), None)
del st.session_state.globalRenames[index]
# Shift remaining widget keys down — ui.py will re-sync them to display
# labels on the next render, so just clear them to force a clean re-seed.
for i in range(index, len(st.session_state.globalRenames)):
st.session_state.pop(_global_rename_key(i), None)
def apply_inline_rename(currFile, raw_sp, new_name):
"""Write a rename from the Rename Speaker tab into speakerRenames and globalRenames."""
new_name = new_name.strip()
token = f"{currFile}: {raw_sp}"
for idx, entry in enumerate(st.session_state.globalRenames):
if token in entry["speakers"]:
entry["speakers"].remove(token)
st.session_state.pop(_global_rename_key(idx), None)
if new_name:
_write_rename(token, new_name)
for idx, entry in enumerate(st.session_state.globalRenames):
if entry["name"] == new_name:
entry["speakers"].append(token)
st.session_state.pop(_global_rename_key(idx), None)
return
st.session_state.globalRenames.append({"name": new_name, "speakers": [token]})
else:
_write_rename(token, "")
st.toast(f"Reverted {raw_sp} to original label")
def on_grename_change(idx, token_display_map=None):
"""Callback for the sidebar rename multiselect at position idx.
token_display_map: dict {display_label -> raw_token} passed from ui.py.
Widget keys hold display labels; entry["speakers"] must hold raw tokens.
"""
# Guard: the entry may have been deleted (e.g. trash button fired just
# before Streamlit re-fired this multiselect callback for the same index).
if idx >= len(st.session_state.globalRenames):
return
grkey = _global_rename_key(idx)
entry = st.session_state.globalRenames[idx]
name = entry["name"]
display_to_raw = token_display_map or {}
raw_to_display = {v: k for k, v in display_to_raw.items()}
prev = list(entry["speakers"]) # raw tokens in data model
# Widget reports display labels — translate back to raw
reported_display = list(st.session_state.get(grkey, []))
reported = [display_to_raw.get(t, t) for t in reported_display]
# Build the set of tokens that are legitimately available for this entry
# right now (not claimed by any OTHER entry).
other_claimed = {
t
for other_idx, other_entry in enumerate(st.session_state.globalRenames)
if other_idx != idx
for t in other_entry["speakers"]
}
# Spurious-empty guard: Streamlit sometimes re-fires this callback with []
# when available_tokens shrinks (e.g. another entry just claimed a token).
# Only treat it as spurious when prev had MORE than 1 token — if prev had
# exactly 1, the user may genuinely be deselecting it, so always let it through.
if not reported and len(prev) > 1:
all_still_valid = all(t not in other_claimed for t in prev)
if all_still_valid:
# Restore widget key as display labels
st.session_state[grkey] = [raw_to_display.get(t, t) for t in prev]
return
prev_set = set(prev)
new_set = set(reported)
added = new_set - prev_set
removed = prev_set - new_set
# Write the new speakers list to the data model
kept = [t for t in prev if t in new_set]
entry["speakers"] = kept + [t for t in added]
# Sync the widget key to match the data model.
# Do NOT pop the key — popping causes Streamlit to re-initialise the widget
# to [] on the next render (because no `default=` is passed), erasing the
# selection visually even though the data model is correct.
# Sync widget key as display labels
st.session_state[grkey] = [raw_to_display.get(t, t) for t in entry["speakers"]]
# Enforce exclusivity: a speaker can only belong to one rename entry at a time.
# Remove the token from any other entry before writing the new name.
for token in added:
for other_idx, other_entry in enumerate(st.session_state.globalRenames):
if other_idx == idx:
continue
if token in other_entry["speakers"]:
other_entry["speakers"].remove(token)
st.session_state[_global_rename_key(other_idx)] = [
raw_to_display.get(t, t) for t in other_entry["speakers"]
]
_write_rename(token, "")
_write_rename(token, name)
# Revert speakerRenames for tokens genuinely removed from this entry
for token in removed:
_write_rename(token, "")
# ---------------------------------------------------------------------------
# File-switch callback
# ---------------------------------------------------------------------------
def updateMultiSelect():
fileName = st.session_state["select_currFile"]
st.session_state.resetResult = True
result = st.session_state.results.get(fileName)
if not result:
return
# Pop category widget keys so they re-seed from categorySelect data
for category in st.session_state.categories:
st.session_state.pop(f"multiselect_{category}", None)
# Pop globalRenames widget keys so they re-seed from entry["speakers"] data
for i in range(len(st.session_state.globalRenames)):
st.session_state.pop(_global_rename_key(i), None)
# ---------------------------------------------------------------------------
# Speaker-clip session-state helpers
# ---------------------------------------------------------------------------
def store_speaker_clips(fname, annotations, waveform, sample_rate):
"""Generate samples & segments and write them into session state."""
clips, segments = utils.build_speaker_clips(annotations, waveform, sample_rate)
st.session_state.speakerClips[fname] = clips
st.session_state.speakerSegments[fname] = segments
st.session_state.speakerWaveforms[fname] = (waveform, sample_rate)
print(f"Generated {len(clips)} speaker samples for {fname}")
def randomize_speaker_clip(file_index, speaker):
"""Replace a speaker's audio sample with a freshly randomized one."""
segs = st.session_state.speakerSegments.get(file_index, {}).get(speaker)
waveform_data = st.session_state.speakerWaveforms.get(file_index)
if not segs or waveform_data is None:
return
waveform, sample_rate = waveform_data
new_clip = utils.get_randomized_clip(waveform, sample_rate, segs)
st.session_state.speakerClips[file_index][speaker] = new_clip
print(f"Randomized sample for {speaker} in {file_index}")
# ---------------------------------------------------------------------------
# Per-file registration helper (keeps Demo / upload code DRY)
# ---------------------------------------------------------------------------
def register_file(fname):
"""Ensure all session-state dicts have an entry for fname."""
st.session_state.results.setdefault(fname, [])
st.session_state.summaries.setdefault(fname, {})
st.session_state.unusedSpeakers.setdefault(fname, [])
# Ensure categorySelect has one list per category (global, not per-file)
while len(st.session_state.categorySelect) < len(st.session_state.categories):
st.session_state.categorySelect.append([])
st.session_state.speakerRenames.setdefault(fname, {})
st.session_state.speakerClips.setdefault(fname, {})
if fname not in st.session_state.file_names:
st.session_state.file_names.append(fname)
# ---------------------------------------------------------------------------
# File loading helpers
# ---------------------------------------------------------------------------
def load_annotation_file(fname, fpath):
"""Load an annotation-only file (.txt / .rttm / .csv) into session state."""
ext = fpath.lower()
if ext.endswith(".txt"):
_, annotations = su.loadAudioTXT(fpath)
elif ext.endswith(".rttm"):
_, annotations = su.loadAudioRTTM(fpath)
elif ext.endswith(".csv"):
_, annotations = su.loadAudioCSV(fpath)
else:
raise ValueError(f"Unsupported annotation format: {fpath}")
totalSeconds = max((s.end for s in annotations.itersegments()), default=0)
st.session_state.results[fname] = (annotations, totalSeconds)
st.session_state.summaries[fname] = {}
st.session_state.unusedSpeakers[fname] = list(annotations.labels())
return annotations, totalSeconds
def load_demo_single(demo_path):
"""Register and load a single RTTM demo file, then run analyze()."""
import time
dname = demo_path.split("/")[-1]
register_file(dname)
st.session_state.file_paths[dname] = demo_path
start_time = time.time()
with st.spinner("Loading Demo Sample"):
load_annotation_file(dname, demo_path)
with st.spinner("Analyzing Demo Data"):
analyze(dname)
st.success(f"Took {time.time() - start_time:.1f}s to analyze the demo file!")
st.session_state.select_currFile = dname
return dname
def load_demo_single_sample(sample_path):
"""Register and load the pre-made short RTTM demo file, then run analyze()."""
import time
dname = sample_path.split("/")[-1]
register_file(dname)
st.session_state.file_paths[dname] = sample_path
start_time = time.time()
with st.spinner("Loading Sample Demo"):
load_annotation_file(dname, sample_path)
with st.spinner("Analyzing Sample Demo Data"):
analyze(dname)
st.success(f"Took {time.time() - start_time:.1f}s to analyze the sample demo!")
st.session_state.select_currFile = dname
return dname
def load_demo_multi(demo_paths):
"""Register and load multiple RTTM demo files."""
for demo_path in demo_paths:
dname = demo_path.split("/")[-1]
register_file(dname)
st.session_state.file_paths[dname] = demo_path
with st.spinner(f"Loading: {dname}"):
load_annotation_file(dname, demo_path)
st.session_state.analyzeAllToggle = True
def run_analysis_loop(file_names, file_paths_dict, pipeline,
enable_denoise, early_cleanup,
gain_window, minimum_gain, maximum_gain,
df_model, df_state, atten_lim_db):
"""Process only new (not yet analyzed) files and populate session state."""
import time
import utils as _utils
start_time = time.time()
# Only process files that haven't been fully analyzed yet.
# A file is considered done only when both results AND summaries are
# populated — load_annotation_file sets results but not summaries, so
# demo/annotation-only files correctly appear in pending until analyze()
# has actually run.
pending = [
fname for fname in file_names
if not (
fname in st.session_state.results
and len(st.session_state.results[fname]) == 2
and st.session_state.summaries.get(fname, {}).get("speakers_dataFrame") is not None
)
]
if not pending:
st.info("All files have already been analyzed.")
st.session_state.analyzeAllToggle = False
return
totalFiles = len(pending)
for i, fname in enumerate(pending):
fpath = file_paths_dict.get(fname, "")
ext = fpath.lower()
if ext.endswith((".txt", ".rttm", ".csv")):
label = ext.rsplit(".", 1)[-1].upper()
with st.spinner(f"Loading {label} {i+1}/{totalFiles}"):
load_annotation_file(fname, fpath)
else:
with st.spinner(f"Processing Audio {i+1}/{totalFiles}"):
annotations, totalSeconds, waveform, sample_rate = _utils.processFile(
fpath, pipeline, enable_denoise, early_cleanup,
gain_window, minimum_gain, maximum_gain,
df_model, df_state, atten_lim_db,
)
st.session_state.results[fname] = (annotations, totalSeconds)
st.session_state.summaries[fname] = {}
st.session_state.unusedSpeakers[fname] = list(annotations.labels())
with st.spinner(f"Generating audio samples {i+1}/{totalFiles}"):
store_speaker_clips(fname, annotations, waveform, sample_rate)
del waveform
with st.spinner(f"Analyzing {i+1}/{totalFiles}"):
analyze(fname)
st.success(f"Analyzed {totalFiles} new file(s) in {time.time() - start_time:.1f}s")
st.session_state.analyzeAllToggle = False
# Rotate uploader key to clear the file uploader widget
st.session_state.uploader_key = st.session_state.get("uploader_key", 0) + 1
def build_table_df(displayDF):
"""Return a display-only copy of displayDF with cosmetic transforms applied:
- Rename 'Resource' -> 'Speaker'
- Drop 'Task' column if present
- Format Start / Finish as HH:MM:SS.cs strings
"""
def _fmt(val):
try:
secs = float(val)
except (TypeError, ValueError):
return str(val)
h = int(secs // 3600)
m = int(secs % 3600 // 60)
s = int(secs % 60)
cs = round((secs % 1) * 100)
return f"{h:02d}:{m:02d}:{s:02d}.{cs:02d}"
df = displayDF.copy()
if "Task" in df.columns:
df = df.drop(columns=["Task"])
if "Start" in df.columns:
df["Start"] = df["Start"].apply(_fmt)
if "Finish" in df.columns:
df["Finish"] = df["Finish"].apply(_fmt)
return df.rename(columns={"Resource": "Speaker"})
# ---------------------------------------------------------------------------
# analyze() — build and cache all DataFrames for one file
# ---------------------------------------------------------------------------
def analyze(inFileName):
"""Compute and store all summary DataFrames for inFileName."""
try:
printV(f"Start analyzing {inFileName}", 4)
st.session_state.resetResult = False
if not (
inFileName in st.session_state.results
and inFileName in st.session_state.summaries
and len(st.session_state.results[inFileName]) > 0
):
return
currAnnotation, currTotalTime = st.session_state.results[inFileName]
speakerNames = currAnnotation.labels()
# categorySelect is global tokens ("fname: SPEAKER_##"); extract raw IDs for this file
prefix = inFileName + ": "
categorySelections = [
[token[len(prefix):] for token in tokens if token.startswith(prefix)]
for tokens in st.session_state.categorySelect
]
printV("Loaded results", 4)
noVoice, oneVoice, multiVoice = su.calcSpeakingTypes(currAnnotation, currTotalTime)
sumNoVoice = su.sumTimes(noVoice)
sumOneVoice = su.sumTimes(oneVoice)
sumMultiVoice = su.sumTimes(multiVoice)
# df3
df3 = utils.build_df3(noVoice, oneVoice, multiVoice)
st.session_state.summaries[inFileName]["df3"] = df3
printV("Set df3", 4)
# df4
df4, nameList, valueList, extraNames, extraValues = utils.build_df4(
speakerNames, categorySelections, st.session_state.categories, currAnnotation
)
st.session_state.summaries[inFileName]["df4"] = df4
printV("Set df4", 4)
# df5
df5 = utils.build_df5(
oneVoice, multiVoice,
sumNoVoice, sumOneVoice, sumMultiVoice,
currTotalTime,
)
st.session_state.summaries[inFileName]["df5"] = df5
printV("Set df5", 4)
# speakers_dataFrame, df2
speakers_dataFrame, speakers_times = su.annotationToDataFrame(currAnnotation)
st.session_state.summaries[inFileName]["speakers_dataFrame"] = speakers_dataFrame
st.session_state.summaries[inFileName]["speakers_times"] = speakers_times
df2 = utils.build_df2(
nameList + extraNames,
valueList + extraValues,
currTotalTime,
)
st.session_state.summaries[inFileName]["df2"] = df2
printV("Set df2", 4)
except Exception as e:
print(f"Error in analyze: {e}")
traceback.print_exc()
st.error(f"Debug - analyze() failed: {e}")