33import pytest
44from causarray .DR_learner import LFC
55
6- # package/tests/test_DR_learner.py
7-
86
97@pytest .fixture
108def sample_data ():
11- np .random .seed (0 )
12- Y = np . random . poisson (5 , (100 , 10 ))
13- W = np . random . normal ( 0 , 1 , (100 , 5 ))
14- A = np . random .binomial (1 , 0.5 , 100 )
15- Y += np . random .poisson (1 , (100 , 10 )) * A [:, None ]
9+ rng = np .random .default_rng (0 )
10+ Y = rng . poisson (5 , (100 , 10 )). astype ( float )
11+ W = rng . standard_normal ( (100 , 5 ))
12+ A = rng .binomial (1 , 0.5 , 100 )
13+ Y += rng .poisson (1 , (100 , 10 )) * A [:, None ]
1614 return Y , W , A
1715
18- def test_LFC_basic (sample_data ):
19- Y , W , A = sample_data
20- result , estimation = LFC (Y , W , A )
21- assert isinstance (result , pd .DataFrame )
22- assert 'tau' in result .columns
23- assert 'std' in result .columns
24- assert 'stat' in result .columns
25- assert 'rej' in result .columns
26- assert 'pvalue' in result .columns
27- assert 'padj' in result .columns
28- assert 'pvalue_emp_null_adj' in result .columns
29- assert 'padj_emp_null_adj' in result .columns
3016
31- def test_LFC_with_offset (sample_data ):
32- Y , W , A = sample_data
33- offset = np .log (np .random .poisson (5 , 100 ))
34- result , estimation = LFC (Y , W , A , offset = offset )
35- assert isinstance (result , pd .DataFrame )
36- assert 'tau' in result .columns
17+ class TestLFCOutputSchema :
18+ def test_output_columns (self , sample_data ):
19+ Y , W , A = sample_data
20+ result , estimation = LFC (Y , W , A )
21+ assert isinstance (result , pd .DataFrame )
22+ for col in ('tau' , 'std' , 'stat' , 'rej' , 'pvalue' , 'padj' ,
23+ 'pvalue_emp_null_adj' , 'padj_emp_null_adj' ):
24+ assert col in result .columns
3725
38- # def test_LFC_with_cross_est(sample_data):
39- # Y, W, A = sample_data
40- # result, estimation = LFC(Y, W, A, cross_est=True)
41- # assert isinstance(result, pd.DataFrame)
42- # assert 'tau' in result.columns
26+ def test_with_offset (self , sample_data ):
27+ Y , W , A = sample_data
28+ rng = np .random .default_rng (1 )
29+ offset = np .log (rng .poisson (5 , 100 ).clip (1 ))
30+ result , estimation = LFC (Y , W , A , offset = offset )
31+ assert isinstance (result , pd .DataFrame )
32+ assert 'tau' in result .columns
4333
44- def test_LFC_with_fdx ( sample_data ):
45- Y , W , A = sample_data
46- result , estimation = LFC (Y , W , A , fdx = True )
47- assert isinstance (result , pd .DataFrame )
48- assert 'tau' in result .columns
34+ def test_with_fdx ( self , sample_data ):
35+ Y , W , A = sample_data
36+ result , estimation = LFC (Y , W , A , fdx = True )
37+ assert isinstance (result , pd .DataFrame )
38+ assert 'tau' in result .columns
4939
50- def test_LFC_with_custom_family ( sample_data ):
51- Y , W , A = sample_data
52- result , estimation = LFC (Y , W , A , family = 'poisson' )
53- assert isinstance (result , pd .DataFrame )
54- assert 'tau' in result .columns
40+ def test_with_custom_family ( self , sample_data ):
41+ Y , W , A = sample_data
42+ result , estimation = LFC (Y , W , A , family = 'poisson' )
43+ assert isinstance (result , pd .DataFrame )
44+ assert 'tau' in result .columns
5545
5646
57- def test_LFC_multi_treatment ():
58- """Test LFC with multiple perturbations (exercises per-perturbation fast path)."""
59- np .random .seed (42 )
60- n , p , a = 200 , 50 , 3
61- W = np .random .normal (0 , 1 , (n , 3 ))
62- # One-hot treatment: first 50 cells = control, rest split among 3 perturbations
63- A = np .zeros ((n , a ))
64- A [50 :100 , 0 ] = 1
65- A [100 :150 , 1 ] = 1
66- A [150 :200 , 2 ] = 1
67- Y = np .random .poisson (5 , (n , p ))
68- # Add treatment effects for perturbation 0
69- Y [50 :100 ] += np .random .poisson (2 , (50 , p ))
47+ class TestLFCMultiTreatment :
48+ def test_multi_treatment (self ):
49+ """LFC with multiple perturbations exercises the per-perturbation fast path."""
50+ np .random .seed (42 )
51+ n , p , a = 200 , 50 , 3
52+ W = np .random .normal (0 , 1 , (n , 3 ))
53+ A = np .zeros ((n , a ))
54+ A [50 :100 , 0 ] = 1
55+ A [100 :150 , 1 ] = 1
56+ A [150 :200 , 2 ] = 1
57+ Y = np .random .poisson (5 , (n , p ))
58+ Y [50 :100 ] += np .random .poisson (2 , (50 , p ))
7059
71- result , estimation = LFC (Y , W , A , family = 'nb' )
72- assert isinstance (result , pd .DataFrame )
73- assert 'trt' in result .columns
74- assert len (result ) == p * a
75- assert result ['tau' ].notna ().all ()
76- assert result ['padj' ].notna ().all ()
60+ result , estimation = LFC (Y , W , A , family = 'nb' )
61+ assert isinstance (result , pd .DataFrame )
62+ assert 'trt' in result .columns
63+ assert len (result ) == p * a
64+ assert result ['tau' ].notna ().all ()
65+ assert result ['padj' ].notna ().all ()
0 commit comments