
    ^j                        d dl mZ d dlmZ d dlmZ d dlmZ d dl	Z
d dlmZ d dlmZ d dlmZmZmZmZ d dlmZ d dlmZ d d	lmZ d d
lmZ  e       Z G d d      Zy)    )annotations)Path)AnyN)Axes)BoxAnnotatorColor
DetectionsLabelAnnotator)Tensor)
DataLoader)box_cxcywh_to_xyxy)
get_loggerc                  d    e Zd ZdZ	 d	 	 	 	 	 	 	 	 	 ddZddZe	 	 	 	 	 	 	 	 	 	 	 	 	 	 d	d       Zy)
DatasetGridSaveraa  Utility for saving 3x3 image grids to visualize augmentation effects.

    Args:
        data_loader: Dataloader of the dataset to sample images from.
        output_dir: Directory where grid images will be saved.
        max_batches: Number of batches to draw samples from.
        dataset_type: Dataset split label, e.g. ``'train'`` or ``'val'``.
    c                v    || _         || _        || _        || _        | j                  j	                  dd       y )NT)parentsexist_ok)data_loader
output_dirmax_batchesdataset_typemkdir)selfr   r   r   r   s        e/var/www/ramen.bs-engineer-server.com/venv/lib/python3.12/site-packages/rfdetr/datasets/save_grids.py__init__zDatasetGridSaver.__init__#   s:     '$&(dT:    c           
        t        j                  g dg d      }t        d      }t        t        j
                  dd      }t        | j                        D ]/  \  }\  }}|| j                  k\  r nt        j                  ddd	
      \  }}|j                  | j                   d|        |j                         }d}	t        t        |j                  |            D ](  \  }	\  }
}|	dk\  r n| j!                  |
|||	   |||       * t#        |	d      D ]  }||   j%                  d        |j'                          t        j(                  | j*                  | j                   d| dz  d       t        j,                          2 t.        j1                  d| j                   d| j*                  j3                                 y)zCreate and save image grids to ``output_dir``.

        Each grid is a 3x3 JPEG containing up to 9 images from a single batch, with bounding boxes and class labels
        drawn on top.
        )g:ܟw g$I$I ggE#)g!:ܟw@gm۶m@grq@)meanstd   )	thicknessg      ?   )
text_color
text_scaletext_padding)   r&   )figsizez dataset, batch r   	   off_batchz	_grid.jpg   )dpizSaved z! grids with augmented images to: N)T	Normalizer   r
   r   BLACK	enumerater   r   pltsubplotssuptitler   flattenziptensors_annotate_and_plotrangeaxistight_layoutsavefigr   closeloggerinforesolve)r   inv_normalizebox_annotatorlabel_annotator	batch_idxsampletargetfigaxessample_indexsingle_imagesingle_targetis                r   	save_gridzDatasetGridSaver.save_grid,   s    A1
 %q1({{
 ,5T5E5E+F 	'I'D,,,Q8<ICLLD--..>ykJK<<>DL?HV^^]cId?e ;;|]1$'' -l1C]Tacr <+ $QU#$ KKT->->,?vi[PY*ZZ`cdIIK+	. 	fT..//PQUQ`Q`QhQhQjPklmr   c                   ddl m} |d   }t        |t              r|j	                         j                         }t        |d         t        |d         }	} ||       }
t        |
t              r,|
j	                         j                         j                         }
|j                  t        j                  |
j                  ddd      dd      dz  j                  t        j                              }t        |d	         dkD  rN|d
   }t        |t              r@|j	                         j                         j                         j                  t              }nt        j                  |t              }|d	   }t        |t              r|j	                         j                         }n|}t        j                  |D cg c]+  }t!        |      }|d   |	z  |d   |z  |d   |	z  |d   |z  g- c}}t        j"                        }t%        ||      }|D cg c]  }t'        |       }}|j)                  ||      }|j)                  |||      }|j+                  |       |j-                  d       yc c}}w c c}w )a|  De-normalize a single image tensor, annotate it with boxes and labels, and plot it on ``ax``.

        Args:
            single_image: Normalized image tensor of shape ``(C, H, W)``.
            single_target: Target dict with keys ``'size'``, ``'boxes'`` (cx,cy,wh normalized), and ``'labels'``.
            ax: Matplotlib axis to plot the annotated image on.
            inv_normalize: Inverse normalization transform to convert the tensor back to pixel values.
            box_annotator: ``BoxAnnotator`` instance for drawing bounding boxes.
            label_annotator: ``LabelAnnotator`` instance for drawing class labels.
        r   )Imagesize   r    g        g      ?   boxeslabels)dtyper"   )xyxyclass_id)scene
detections)rW   rX   rS   r)   N)PILrN   
isinstancer   detachcpuintnumpy	fromarraynpclip	transposeastypeuint8lenasarrayr   float32r	   strannotateimshowr9   )rI   rJ   axr@   rA   rB   PILImageresized_sizehwde_normalized_imgrW   labels_tensor	class_idsrR   
boxes_iterboxbrU   rX   crS   s                         r   r7   z#DatasetGridSaver._annotate_and_plotV   s%   & 	*$V,lF+'..0446L<?#Sa%91),7'0 1 8 8 : > > @ F F H""BGG,=,G,G1a,PRUWZ$[^a$a#i#ijljrjr#st}W%&*)(3M-0)002668>>@GGL	JJ}C@	!'*E%("\\^//1
"
::EOscZlmpZqTU!A$(AaD1HadQh!q9sjjD $	BJ&/0c!f0F0!**:*NE#,,5ZX^,_E
		%
 t 1s   =0I4
I:N)r"   train)
r   r   r   r   r   r]   r   rh   returnNone)rx   ry   )rI   r   rJ   zdict[str, Any]rk   r   r@   zT.NormalizerA   r   rB   r
   rx   ry   )__name__
__module____qualname____doc__r   rL   staticmethodr7    r   r   r   r      s     dk;%;37;FI;]`;	;(nT 55%5 5 #	5
 $5 (5 
5 5r   r   )
__future__r   pathlibr   typingr   matplotlib.pyplotpyplotr1   r^   r`   torchvision.transforms
transformsr-   matplotlib.axesr   supervisionr   r   r	   r
   torchr   torch.utils.datar   rfdetr.util.box_opsr   rfdetr.util.loggerr   r=   r   r   r   r   <module>r      sA    #     "   G G  ' 2 )	s sr   