| """Deterministic contract tests; real audio validation is documented separately.""" |
| from types import SimpleNamespace |
| from pathlib import Path |
| import numpy as np |
| import pytest |
| import torch |
| from streampa.pipeline import StreamPASession, load_scorer |
| from streampa import frontend as fe |
|
|
| class ScheduleSession(StreamPASession): |
| def _step(self,commit_samples,end_samples,is_final): |
| row=dict(commit=commit_samples,end=end_samples,final=is_final,samples=self.audio[:end_samples].copy()) |
| self.last_result=row |
| return row |
|
|
| def session(): |
| return ScheduleSession(SimpleNamespace(assets={'max_audio_sec':30,'chunk_sec':.64,'right_context_sec':.16})) |
|
|
| def test_received_boundary_and_short_eof(): |
| s=session() |
| assert s.append(np.zeros(12799,np.float32))==[] |
| r=s.append(np.ones(1,np.float32))[0] |
| assert r['commit']==10240 and r['end']==12800 and not r['final'] |
| s.append(np.ones(100,np.float32)) |
| r=s.finish() |
| assert r['end']==12900 and r['commit']==12900 and r['final'] |
|
|
| def test_reset_and_new_session_do_not_share_state(): |
| a=session(); b=session() |
| a.state=torch.ones(1,64); a.commit_state={'committed_slots':[0]} |
| a.append(np.ones(100)); a.reset() |
| assert a.state is None and a.commit_state is None and len(a.audio)==0 |
| assert b.state is None and len(b.audio)==0 |
|
|
| def test_suffix_is_not_in_current_input(): |
| prefix=np.arange(12800,dtype=np.float32) |
| a=session(); b=session() |
| x=a.append(np.concatenate([prefix,np.ones(10000)]))[0] |
| y=b.append(np.concatenate([prefix,-np.ones(10000)]))[0] |
| assert np.array_equal(x['samples'],y['samples']) |
|
|
| def test_closed_empty_and_invalid_input(): |
| s=session(); assert s.finish()['end']==0 |
| with pytest.raises(RuntimeError): s.append(np.zeros(1)) |
| s.reset() |
| with pytest.raises(ValueError): s.append(np.array([np.nan])) |
| with pytest.raises(ValueError): s.append(np.zeros((2,2))) |
| with pytest.raises(ValueError): s.append(np.zeros(480001)) |
| assert len(s.audio)==0 |
|
|
| def test_normalization_and_gap_mass(): |
| assert fe.normalize_phone('AH0')=='AH' |
| assert fe.normalize_word('Hello!')=='HELLO' |
| post=np.array([[.6,.1,.3],[.2,.6,.2]],np.float32) |
| pcn={'cn_post':post,'eps_index':2} |
| fe.validate_pcn(pcn) |
| stats,_=fe.pcn_stats(post,[0,1],[],2) |
| assert np.allclose(stats[:,0],[.3,.2]) |
| with pytest.raises(ValueError): fe.validate_pcn({'cn_post':post*0,'eps_index':2}) |
|
|
| def test_partial_commit_and_legacy_persistence(): |
| pcn={'top_phone_ids':[1,2],'slot_times':[(0,.3),(.3,.8)]} |
| ids=np.array([0,1],np.int32) |
| times=[{'word':'A','start':0,'end':.3},{'word':'B','start':.3,'end':.8}] |
| cum,new,_,_,state=fe.build_stateful_commit_masks(None,pcn,ids,times,.64,.8,False,commit_rule='legacy') |
| assert new.tolist()==[1.,0.] and cum.tolist()==[1.,0.] |
| cum,new,_,_,state=fe.build_stateful_commit_masks(state,pcn,ids,times,.8,.8,True,commit_rule='legacy') |
| assert new.tolist()==[0.,1.] and cum.tolist()==[1.,1.] |
| _,new,_,_,_=fe.build_stateful_commit_masks(state,pcn,ids,times,1,1,True,commit_rule='legacy') |
| assert new.sum()==0 |
|
|
| def test_no_new_words_keeps_state_but_visible_path_remains_active(): |
| torch.set_num_threads(2); torch.manual_seed(1337) |
| model,c=load_scorer(Path(__file__).resolve().parents[1]) |
| n=c['seq_len']; d=c['phone_dim'] |
| batch=dict(cn_post=torch.softmax(torch.randn(1,n,d),-1),cn_stats=torch.randn(1,n,5), |
| acoustic_post=torch.softmax(torch.randn(1,n,d),-1),acoustic_stats=torch.randn(1,n,4), |
| prosody=torch.randn(1,c['prosody_dim']),slot_prosody=torch.randn(1,n,c['slot_prosody_dim']), |
| visible_len=torch.tensor([4]),cumulative_commit_mask=torch.zeros(1,n),new_commit_mask=torch.zeros(1,n), |
| word_ids=torch.arange(n)[None]) |
| state=torch.randn(1,1,64) |
| with torch.inference_mode(): |
| first=model(**batch,prev_state=state) |
| batch['cn_stats'][:,:4]+=2 |
| second=model(**batch,prev_state=state) |
| assert torch.equal(first['next_state'],state) and torch.equal(second['next_state'],state) |
| assert first['new_word_mask'].sum()==0 |
| assert not torch.equal(first['utt_scores'],second['utt_scores']) |
|
|