@@ -456,7 +456,6 @@ def __init__(
456456 skip_fields : List [
457457 str
458458 ] = [], # Data fields to exclude during loading to save memory
459- n_parallel : int = 20 , # Number of parallel processes for segmentation
460459 empirical_data_fn : (
461460 str | None
462461 ) = None , # Path to empirical distribution data file for phase-to-angle mapping
@@ -477,7 +476,6 @@ def __init__(
477476 ):
478477 # Store configuration parameters
479478 self .yaml_fn = yaml_fn
480- self .n_parallel = n_parallel
481479 self .nthetas = nthetas # Number of angles to discretize space for beamforming
482480 self .target_ntheta = self .nthetas if target_ntheta is None else target_ntheta
483481
@@ -645,14 +643,14 @@ def check_for_new_data(self):
645643
646644 def write_to_idx (self , idx , ridx , raw ):
647645 # this is the heavy lifting of processing, do it on this process
648- rendered_data = self .render_session (idx , ridx , raw )
646+ rendered_data = self .render_session (ridx , raw )
649647
650648 self .incoming_queue .put ((idx , ridx , rendered_data ))
651649
652650 # with self.condition:
653651 # self.condition.notify_all()
654652
655- def render_session (self , idx , ridx , data ):
653+ def render_session (self , ridx , data ):
656654 snapshot_idxs = [0 ] # which snapshots to get
657655
658656 data ["rx_wavelength_spacing" ] = torch .tensor (self .rx_wavelength_spacing )
@@ -677,26 +675,29 @@ def render_session(self, idx, ridx, data):
677675
678676 if "signal_matrix" not in self .skip_fields :
679677 # WARNGING this does not respect flipping!
678+ # signal matrix ~ 1,1,2,524288
680679 abs_signal = data ["signal_matrix" ].abs ().to (torch .float32 )
681680 assert data ["signal_matrix" ].shape [0 ] == 1
682681 pd = torch_get_phase_diff (data ["signal_matrix" ][0 ]).to (torch .float32 )
683682 data ["abs_signal_and_phase_diff" ] = torch .concatenate (
684683 [abs_signal , pd [None , :, None ]], dim = 2
685684 )
686685
687- data ["rx_pos_mm" ] = torch .vstack (
688- [
689- data ["rx_pos_x_mm" ],
690- data ["rx_pos_y_mm" ],
691- ]
692- ).T
686+ # data["rx_pos_mm"] = torch.vstack(
687+ # [
688+ # data["rx_pos_x_mm"], # size = [1]
689+ # data["rx_pos_y_mm"], # size = [1]
690+ # ]
691+ # ).T # torch.Size([1, 2])
693692
694- data ["tx_pos_mm" ] = torch .vstack (
695- [
696- data ["tx_pos_x_mm" ],
697- data ["tx_pos_y_mm" ],
698- ]
699- ).T
693+ # data["tx_pos_mm"] = torch.vstack(
694+ # [
695+ # data["tx_pos_x_mm"], # size = [1]
696+ # data["tx_pos_y_mm"], # size = [1]
697+ # ]
698+ # ).T # torch.Size([1, 2])
699+
700+ data ["rx_pos_mm" ] = data ["tx_pos_mm" ] = torch .ones (1 , 2 ) * torch .nan
700701
701702 data ["rx_pos_xy" ] = (
702703 data ["rx_pos_mm" ][snapshot_idxs ].unsqueeze (0 ) / self .distance_normalization
@@ -705,9 +706,13 @@ def render_session(self, idx, ridx, data):
705706 data ["tx_pos_xy" ] = (
706707 data ["tx_pos_mm" ][snapshot_idxs ].unsqueeze (0 ) / self .distance_normalization
707708 )
708- breakpoint ()
709+
710+ signal_matrix = data ["signal_matrix" ][0 ][0 ]
711+ if isinstance (signal_matrix , torch .Tensor ):
712+ signal_matrix = signal_matrix .numpy ()
713+
709714 segmentation = segment_session (
710- data [ " signal_matrix" ][ 0 ][ 0 ]. numpy () ,
715+ signal_matrix ,
711716 gpu = self .gpu ,
712717 skip_beamformer = False ,
713718 skip_detrend = self .skip_detrend ,
0 commit comments