@@ -181,3 +181,91 @@ def score_transports_and_targets_combinations(all_expr, all_meta):
181181 scores ["scores" ] = scores ["scores" ].astype (float )
182182
183183 return scores
184+
185+ def monge_get_source_target_transport (
186+ trainer : "MongeGapTrainer" ,
187+ datamodule : "AbstractDataModule" ,
188+ n_samples = 1 ,
189+ target = True ,
190+ source = True ,
191+ transport = True ,
192+ batch_size : int = None ,
193+ ):
194+ if batch_size is None :
195+ batch_size = datamodule .batch_size
196+
197+ if datamodule .split [2 ] > 0 :
198+ print ("Evaluating on test set" )
199+ split_type = "test"
200+ elif datamodule .split [1 ] > 0 :
201+ print ("Evaluating on validation set" )
202+ split_type = "valid"
203+ else :
204+ print ("Evaluating on training set" )
205+ split_type = "train"
206+ all_expr = []
207+ all_meta = []
208+ for i in range (n_samples ):
209+ if split_type == "valid" :
210+ sel_target_cells = random .sample (
211+ datamodule .target_valid_cells .tolist (),
212+ batch_size ,
213+ )
214+ sel_control_cells = random .sample (
215+ datamodule .control_valid_cells .tolist (),
216+ batch_size ,
217+ )
218+
219+ elif split_type == "test" :
220+ sel_target_cells = random .sample (
221+ datamodule .target_test_cells .tolist (), batch_size
222+ )
223+ sel_control_cells = random .sample (
224+ datamodule .control_test_cells .tolist (),
225+ batch_size ,
226+ )
227+ elif split_type == "train" :
228+ sel_target_cells = random .sample (
229+ datamodule .target_train_cells .tolist (),
230+ batch_size ,
231+ )
232+ sel_control_cells = random .sample (
233+ datamodule .control_train_cells .tolist (),
234+ batch_size ,
235+ )
236+
237+ cond_expr = datamodule .adata [sel_target_cells ].X
238+ cond_meta = datamodule .adata .obs .loc [sel_target_cells , :]
239+
240+ source_expr = datamodule .adata [sel_control_cells ].X
241+ source_meta = datamodule .adata .obs .loc [sel_control_cells , :]
242+
243+ if target :
244+ cond_meta ["sample_n" ] = i
245+ cond_meta ["dtype" ] = "target"
246+ all_expr .append (pd .DataFrame (cond_expr , columns = datamodule .adata .var_names ))
247+ all_meta .append (cond_meta )
248+
249+ if source :
250+ source_meta ["dtype" ] = "source"
251+ source_meta ["sample_n" ] = i
252+
253+ all_meta .append (source_meta )
254+ all_expr .append (
255+ pd .DataFrame (source_expr , columns = datamodule .adata .var_names )
256+ )
257+
258+ if transport :
259+ trans = trainer .transport (source_expr )
260+ trans = datamodule .decoder (trans )
261+ trans_meta = cond_meta .copy ()
262+ trans_meta ["dtype" ] = "transport"
263+ trans_meta ["sample_n" ] = i
264+
265+ all_expr .append (pd .DataFrame (trans , columns = datamodule .adata .var_names ))
266+ all_meta .append (trans_meta )
267+
268+ all_expr = pd .concat (all_expr ).reset_index (drop = True )
269+ all_meta = pd .concat (all_meta ).reset_index (drop = True )
270+
271+ return all_expr , all_meta
0 commit comments