Visual analysis framework for model verification based on interpretable data slices
Through a visual analysis workflow based on interpretable data slicing, the problem of model verification in the prior art requires a large amount of metadata and manpower, and efficient model verification and performance improvement without metadata are achieved.
Patent Information
- Application Number
- CN202411749446.4
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Priority Date
- 2023-11-30
- Filing Date
- 2024-12-02
- Publication Date
- 2025-05-30
AI Technical Summary
The prior art requires a lot of metadata and manpower when performing machine learning model verification, and it is difficult to explain the root cause of poorly performed data slicing.
Using a visual analysis workflow based on interpretable data slicing, a feature vector is generated through pixel attributes, cluster data slicing, and a graphical user interface for data slicing mosaics allows users to annotate and modify training data to improve model performance.
Meaningful data slices are generated without additional metadata, significantly reducing the time and labor costs of model validation and providing in-depth insights into model behavior to help improve model performance.
Smart Images

Figure CN120069099A_ABST
Abstract
Description
Technical Field
[0001] The devices and methods disclosed in this document relate to visual analysis for machine learning, and more particularly, to a visual analysis framework for model validation based on interpretable data slices. Background Art
[0002] Unless otherwise indicated herein, the materials described in this section are not admitted to be prior art by virtue of inclusion in this section.
[0003] Machine learning (ML) model validation refers to the process of determining and understanding when and why a model succeeds or fails. As machine learning techniques become more prevalent in various fields, model validation has become increasingly important for providing greater model transparency, accountability, and accuracy. One such method for machine learning model validation is data slice finding, which seeks to identify specific instances or subsets of data, i.e., data slices, for which the model performs poorly compared to the entire dataset. These data slices are typically defined by tabular features or additional metadata, such as when and where the data was collected. Problematic data slices may come from underrepresented feature sets, such as a water background that only shows water birds and not a general forest background, or from biased samples, such as a subset of facial images that consists mainly of males.
[0004] Once problematic data slices are identified, model developers can attempt to improve the performance of the model using two strategies. First, the developers can simply retrain or fine-tune the model by focusing on these data subgroups to improve the overall reliability of the model. Second, the developers can attempt to use explainable artificial intelligence (XAI) techniques to identify the underlying root causes that affect the model's accuracy and manipulate the internal parameters of the model to improve its performance. However, despite the great benefits of data slice finding and XAI techniques in explaining problematic data slices, there are still challenges that need to be addressed to conduct thorough model validation.
[0005] Existing data slice finding methods require a large amount of metadata associated with the validation dataset, and collecting such metadata typically involves a significant amount of human effort or additional machine learning components, which is expensive. The job of the slice finding method is to group data into subsets that share common characteristics. To make the data groups interpretable by humans, additional metadata is needed. While some current techniques use tabular metadata (labels) to calculate slices, other techniques apply multimodal pre-trained vision-language models to generate these labels. When no metadata or suitable pre-trained model is available, understanding the data subgroups produced by slice finding becomes a time-consuming and challenging task. In particular, users must individually examine multiple data samples of the data slices in order to formulate hypotheses and understand the details of the slices.
[0006] Using XAI techniques to uncover the root causes of model failures on troublesome data slices also requires considerable effort. In particular, there are several distinct challenges. First, most XAI techniques are instance-based, and it is practically infeasible to draw broad conclusions across all instances. Additionally, summarizing failure patterns and obtaining actionable insights without in-depth investigation is a challenge. Summary of the Invention
[0007] Disclosed is a method for validating a vision model. The vision model is configured to receive an input image and provide at least one output label. The method includes storing a plurality of validation images in a memory. The method further includes using a processor to determine, using the vision model, a plurality of first feature vectors, the plurality of first feature vectors including a respective first feature vector for each respective validation image of the plurality of validation images. The method further includes using the processor to identify, based on the plurality of first feature vectors, multiple sets of validation images from the plurality of validation images. The method further includes displaying, on a display screen, a visualization of at least one set of validation images from the multiple sets of validation images.
[0008] Disclosed is a method for validating a machine learning model. The machine learning model is configured to receive an input data sample and provide at least one output label. The method includes storing a plurality of validation data samples in a memory. The method further includes using a processor to determine, using the machine learning model, a plurality of first feature vectors, the plurality of first feature vectors including a respective first feature vector for each respective validation data sample of the plurality of validation data samples. The method further includes using the processor to identify, based on the plurality of first feature vectors, multiple sets of validation data samples from the plurality of validation data samples. The method further includes displaying, on a display screen, a visualization of at least one set of validation data samples from the multiple sets of validation data samples.
[0009] A system for validating a machine learning model is disclosed. The machine learning model is configured to receive input data samples and provide at least one output label. The system includes a display screen configured to display a graphical user interface. The system also includes at least one memory configured to store program instructions of the machine learning model, a plurality of validation data samples, and a plurality of training data samples. The system also includes at least one processor. The at least one processor is configured to use the machine learning model to determine a plurality of first feature vectors, the plurality of first feature vectors including a respective first feature vector for each respective validation data sample of the plurality of validation data samples. The at least one processor is also configured to identify multiple sets of validation data samples from the plurality of validation data samples. The at least one processor is also configured to display a visualization of at least one set of validation data samples from the multiple sets of validation data samples within the graphical user interface. The at least one processor is also configured to annotate at least one corresponding set of validation data samples from the multiple sets of validation data samples with at least one annotation based on user input. The at least one processor is also configured to modify training data samples among the plurality of training data samples based on the at least one annotation. The at least one processor is also configured to retrain the machine learning model using the modified plurality of training data samples. BRIEF DESCRIPTION OF THE DRAWINGS
[0010] The foregoing aspects and other features of the method and system are explained in the following description in conjunction with the accompanying drawings.
[0011] Figure 1 An overview of the workflow for visual analysis using interpretable data slices is presented, which does not require any metadata.
[0012] Figure 2 An exemplary embodiment of a computing device that can be used for machine learning model analysis and validation based on data slices is shown.
[0013] Figure 3 A flowchart of a method for analyzing and validating the performance of a machine learning model is shown.
[0014] Figure 4 An exemplary workflow for determining the feature vectors of each validation data sample is shown.
[0015] Figure 5A and Figure 5B An exemplary graphical user interface for visualizing and analyzing data slices is shown.
[0016] Figure 6 An overview of the workflow for generating a data slice mosaic for a visual model is presented.
[0017] Figure 7 An exemplary workflow for training images to corrupt false features with noise is presented.
[0018] Figure 8 Shows common usage patterns of the system and method.
[0019] Figure 9 Shows some exemplary spuriousness issues detected during the operation of a hair color classifier model.
[0020] Figure 10A and Figure 10B Shows a further exemplary graphical user interface for visualizing analytical data slices.
[0021] Figure 11 Shows some exemplary spuriousness issues detected during the operation of a bird category classifier model.
[0022] Figure 12 Shows a table that outlines a quantitative assessment of the overall performance of a hair color classification model.
[0023] Figure 13 Shows a table that outlines a quantitative assessment of the overall performance of a bird category classification model.
[0024] Figure 14 Shows a qualitative assessment leveraging pixel attributes from GradCAM. Detailed Description
[0025] To facilitate understanding of the principles of the present disclosure, reference will now be made to the embodiments shown in the accompanying drawings and described in the following written description. It should be understood that this does not thereby limit the scope of the present disclosure. It should also be understood that the present disclosure includes any changes and modifications to the shown embodiments, and includes further applications of the principles of the present disclosure that would typically occur to those skilled in the art to which the present disclosure pertains.
[0026] Overview
[0027] Data slice finding is an emerging technique for evaluating machine learning models by finding subgroups with poor performance in a target dataset. These subgroups are typically defined by a feature set or metadata. Although data slice finding is useful, it poses two challenges for unstructured image data: (1) data slice finding typically requires additional metadata, and providing this metadata is labor-intensive and costly, and (2) data slice finding typically requires a great deal of effort to explain the root cause of underperforming data slices. To address these challenges, a novel human-in-the-loop visual analysis workflow for machine learning model validation based on data slices is disclosed.
[0028] Figure 1An overview of a workflow 10 for visual analysis using interpretable data slices is presented, which does not require any metadata. The workflow 10 enables a user to validate and improve a machine learning model 12 that has been trained using a dataset 14. In at least some embodiments, the machine learning model 12 is a visual model (e.g., a hair color classification model), and the dataset 14 includes a large number of images (e.g., images of people with hair). The workflow 10 generally includes three phases: an interpretable data slice finding phase 20, a slice overview and annotation phase 30, and a slice error mitigation phase 40.
[0029] In the interpretable data slice finding phase 20, the workflow 10 employs pixel attributes to create interpretable features of the dataset 14. In one embodiment, the Gradient-weighted Class Activation Mapping (GradCAM) method is used to generate a feature vector by extracting model attributes in the latent space. In some embodiments, the extracted features are weighted to form a weighted feature vector. Finally, the weighted feature vectors are clustered to identify data slices. In this way, the dataset 14 can be sliced based on the intrinsic behavior of the machine learning model 12.
[0030] Next, in the slice overview and annotation phase 30, the workflow 10 transforms the generated features into one or more visualizations, including a "data slice mosaic" and other visual overviews that clarify the model inference for a given data slice. With the help of data slice mosaics and spurious (label) propagation, a user can identify and annotate slice error types, such as core / spurious correlations and noisy (incorrect) labels. In this way, the workflow 10 adopts a human-in-the-loop approach to rank and annotate problematic data slices based on common error types.
[0031] Finally, in the slice error mitigation phase 40, the workflow 10 utilizes the annotations and user-verified spuriousness to mitigate slice errors in the machine learning model 12. Error mitigation techniques include relabeling data in the dataset 14 when necessary and applying the Core Risk Minimization (CoRM) technique if a spurious correlation is detected. In this way, by identifying and correcting errors on the model or data side, the annotated slices can be used to enhance the performance of the machine learning model 12.
[0032] Compared with conventional visual analysis and model verification techniques, the visual analysis workflow 10 offers many advantages. In particular, the visual analysis workflow 10 is advantageously able to find interpretable slices in the data based on interpretable features extracted from pixel attributes without any additional metadata such as text annotations or cross-model embeddings. The visual analysis workflow 10 can advantageously identify key model issues, including spurious features and mislabeled data. The visual analysis workflow 10 advantageously provides a graphical user interface with data slice mosaics that enables users to quickly browse data slice summaries to gain insights into model behavior. In this way, users can quickly understand the content of the data slices, reducing the need to examine individual data samples. Finally, the visual analysis workflow 10 closes the loop by allowing domain experts to mitigate model issues using the resulting insight summaries and state-of-the-art neural network regularization techniques.
[0033] Design Goals
[0034] The visual analysis workflow 10 is a novel method for metadata-free verification of machine learning models based on slice-level analysis. To facilitate a comprehensive understanding of the workflow 10, the design goals informing the development and design of the workflow 10 are discussed.
[0035] Since slice finding can facilitate the comprehensive evaluation of machine learning models, it has recently received considerable attention. However, a key gap has been recognized in slice-based model verification systems for machine learning models, especially visual models. This task is particularly difficult for visual models because the data is not structured. More specifically, these systems require metadata or language-visual models to generate meaningful data slices, which is not always feasible. When metadata is not available, obtaining it can be a laborious annotation process that requires significant human and financial resources. Additionally, the language-visual models used are typically trained on general-purpose datasets that may not be sufficient to represent a specific domain. This problem raises doubts about the adaptability of these models in a customized domain. Based on these observations, the following design goals are compiled for a metadata-free slice-driven model verification workflow:
[0036] Design Goal 1: Metadata-free and model-free. It is advantageous if the data slice finding method of the workflow 10 can generate meaningful data slices without any metadata or visual-language model. The workflow 10 should be able to identify data slices that share important patterns.
[0037] Design Goal 2: Interpretability. It is advantageous if the workflow 10 supports the meaningful interpretation of the identified data slices. Compared with the previous metadata-free slice finding methods used for visual tasks, the workflow 10 should enable the user to clearly understand why specific data samples are grouped in a slice.
[0038] Design Goal 3: Slice Overview. It is advantageous if the workflow 10 minimizes or eliminates any need for individual inspection of each image in the data slice. The workflow 10 should provide an overview of the slices to reduce the time required for analysis, especially when there are a large number of data slices.
[0039] Design Goal 4: Actionable Insights. It is advantageous if the workflow 10 provides actionable insights by allowing the user to annotate and export the data slices. This allows for further expert interaction, such as relabeling or model retraining, thereby improving model performance.
[0040] Based on the description herein, it will be understood that the disclosed visual analysis methods and systems achieve each of these four design goals.
[0041] Exemplary Hardware Embodiment
[0042] Figure 2 An exemplary embodiment of a computing device 100 that can be used for machine learning model analysis and validation based on data slices without metadata is shown. The computing device 100 includes a processor 110, a memory 120, a display screen 130, a user interface 140, and at least one network communication module 150. It will be understood that the illustrated embodiment of the computing device 100 is merely an exemplary embodiment and represents only any one of the various ways or configurations of a server, a desktop computer, a laptop computer, a mobile phone, a tablet computer, or any other computing device that operates in the manner described herein. In at least some embodiments, the computing device 100 communicates with a database 102, which can be hosted by another device or stored in the memory 120 of the computing device 100 itself.
[0043] The processor 110 is configured to execute instructions to operate the computing device 100, thereby enabling the implementation of the features, functions, characteristics, etc. described herein. To this end, the processor 110 is operably connected to the memory 120, the display screen 130, and the network communication module 150. The processor 110 generally includes one or more processors, which can operate in parallel or otherwise cooperate with each other. Those of ordinary skill in the art will recognize that a "processor" includes any hardware system, hardware mechanism, or hardware component that processes data, signals, or other information. Therefore, the processor 110 can include a system having a central processing unit, a graphics processing unit, multiple processing units, dedicated circuits for implementing functions, programmable logic, or other processing systems.
[0044] The memory 120 is configured to store data and program instructions that, when executed by the processor 110, enable the computing device 100 to perform various operations described herein. The memory 120 can be any type of device capable of storing information accessible to the processor 110, such as a memory card, ROM, RAM, hard disk drive, magnetic disk, flash memory, or any of the various other computer-readable media used as data storage devices, as will be recognized by those of ordinary skill in the art.
[0045] The display screen 130 can include any of various known types of displays configured to display a graphical user interface, such as an LCD or OLED screen. The user interface 140 can include various interfaces for operating the computing device 100, such as buttons, switches, keyboards or other keypads, speakers, and microphones. Alternatively or additionally, the display screen 130 can include a touch screen configured to receive touch input from the user.
[0046] The network communication module 150 can include one or more transceivers, modems, processors, memories, oscillators, antennas, or other hardware conventionally included in a communication module to enable communication with various other devices. Specifically, the network communication module 150 generally includes an Ethernet adapter or a module configured to enable communication with a wired or wireless network and / or a router (not shown) configured to enable communication with various other devices. Additionally, the network communication module 150 can include a module (not shown), and one or more cellular modems configured to communicate with a wireless telephone network.
[0047] The memory 120 stores the program instructions of the data slice visualization application 160. Additionally, the memory 120 stores the program instructions and parameters (e.g., kernel weights, model coefficients, etc.) of the machine learning model 12. In at least some embodiments, the database 102 stores a data set 170 including a plurality of data samples (e.g., images). The data samples may include training and validation pairs, each training and validation pair including respective data samples, such as images, and at least one label associated with the respective data sample, such as a hair color label.
[0048] Method for Visualizing Data Slices
[0049] Various operations and processes are described below for operating the computing device 100 to provide machine learning model analysis and validation based on data slices without metadata. In these descriptions, a statement that a method, processor, and / or system is performing a certain task or function refers to a controller or processor (e.g., the processor 110 of the computing device 100) executing programming instructions stored in a non-transitory computer-readable storage medium (e.g., the memory 120 of the computing device 100) operably connected to the controller or processor to manipulate data or operate one or more components in the computing device 100 or the database 102 to perform the task or function. Additionally, the steps of the method can be executed in any feasible chronological order, regardless of the order shown in the figures or the order in which the steps are described.
[0050] Figure 3 A flowchart of a method 200 for analyzing and validating machine learning model performance is shown. The method 200 advantageously enables a user to interpret and understand data slices in a training or validation data set without any additional metadata. The method 200 advantageously enables a user to identify key model issues in the model and the data set, including spurious features and mislabeled data. The method 200 advantageously provides a graphical user interface with intuitive data slice mosaics that enables the user to quickly browse data slice summaries to gain insights into model behavior. In this way, the user can quickly understand the content of the data slices, reducing the need to examine individual data samples. Finally, the method 200 enables the user to use the resulting insight summaries and state-of-the-art neural network regularization techniques to mitigate model problems.
[0051] Method 200 begins with providing a machine learning model that has been trained using a training data set (block 210). Specifically, processor 110 receives and / or memory 120 stores program instructions and parameters (e.g., kernel weights, model coefficients, etc.) of machine learning model 12, which is the object of analysis and verification. As used herein, the term "machine learning model" refers to a system or program instructions and / or data set configured to implement an algorithm, process, or mathematical model (e.g., a neural network), which algorithm, process, or mathematical model predicts or otherwise provides a desired output based on a given input. It will be understood that, generally, many or most of the parameters of a machine learning model are not explicitly programmed, and in the traditional sense, a machine learning model is not explicitly designed to follow specific rules to provide a desired output for a given input. Instead, a corpus of training data is provided to the machine learning model, and the machine learning model identifies or "learns" patterns and statistical relationships in the data from the corpus, which patterns and statistical relationships are generalized to make predictions or otherwise provide an output regarding new data inputs. The results of the training process are embodied in a plurality of learned parameters (e.g., kernel weights, model coefficients, etc.), which are used in various components of the machine learning model to perform various operations or functions.
[0052] In some embodiments, machine learning model 12 is a vision model, such as an image classification model. In one example, machine learning model 12 is a vision model configured to receive an image and output a hair color classification (e.g., non-gray hair, gray hair) of at least one person represented in the image. Similarly, in another example, machine learning model 12 is a vision model configured to receive an image and output a bird category classification of at least one bird represented in the image. In other embodiments, machine learning model 12 may include other types of vision models, such as an object detection model or an image segmentation model.
[0053] In at least some embodiments, machine learning model 12 is a convolutional neural network (CNN) model. It will be understood that a CNN is a feedforward neural network that includes multiple convolutional layers. A traditional convolutional layer receives an input and applies one or more convolutional filters to the input. A convolutional filter (also referred to as a kernel) is a weight matrix (also referred to as a parameter or filter value) that is applied to individual blocks of an input matrix such that the weight matrix is convolved over the input matrix to provide an output matrix. The various layers and filters of a CNN are used to detect various "features" of the input. However, it should be understood that machine learning model 12 may include any architecture of a neural network model, such as a transformer-based architecture (e.g., vision transformer) or a recurrent neural network-based architecture.
[0054] Finally, the processor 110 receives and / or the database 102 stores a data set 170 including a plurality of data samples. The plurality of data samples includes a plurality of validation data samples and a plurality of training data samples for training the machine learning model 12. Generally, the plurality of validation data samples includes data samples not used for training the machine learning model 12. The plurality of data samples may include data-label pairs, each data-label pair including a respective data sample and at least one label associated with the respective data sample. In the case of a vision model, each data-label pair includes an image and a label (e.g., a hair color classification label, an object detection bounding box, an image segmentation label, etc.).
[0055] Method 200 continues with generating interpretable features of the validation data set (block 220). Specifically, the processor 110 determines a plurality of weighted feature vectors F W , which includes respective weighted feature vectors F W of each respective validation data sample among the plurality of validation data samples. Further, in some embodiments, the processor 110 further determines a plurality of top attribute feature vectors F T , which includes respective top attribute feature vectors F T of each respective validation data sample among the plurality of validation data samples.
[0056] In some embodiments where the machine learning model 12 is a neural network with multiple layers, the processor 110 extracts a respective intermediate output F of at least one intermediate neural network layer of the machine learning model 12 by inputting the respective validation data sample into the machine learning model 12, and determines the respective weighted feature vector F W based on the respective intermediate output, to generate each weighted feature vector F W . More specifically, in some embodiments, the processor 110 determines a weight matrix W based on the model gradients of the machine learning model 12, e.g., using gradient-based class activation mapping (CAM) methods such as GradCAM, GradCAM++, and SmoothCAM. Next, the processor 110 determines the respective weighted feature vector F W by using the weight matrix W to determine the weighted average of the respective intermediate output F along at least one dimension. Similarly, the processor 110 determines each respective top attribute feature vector F T by using the weight matrix W to determine the maximum weighted value along at least one dimension in the respective intermediate output F.
[0057] Figure 4 Illustrates determining the respective feature vectors F W and F TExemplary workflow 300. The processor 110 inputs corresponding verification data samples, particularly the image 302, into the machine learning model 12. In the illustrated embodiment, the machine learning model 12 is a CNN model that has a set of convolutional layers followed by fully connected layers. The processor 110 extracts the latent space feature vector F from the intermediate output of the machine learning model 12. In one embodiment, the feature vector F is the output of the final convolutional layer before the fully connected layer of the machine learning model 12. The resulting feature vector F has a dimension of m×n×d.
[0058] Next, the processor 110 generates a weight matrix W with a dimension of m×n using the model gradient. More specifically, the processor 110 uses the GradCAM method to calculate the GradCAM pixel attributes for the image and also extracts its internal weight matrix W. In some embodiments, the processor 110 upsamples the matrix W to the image size of the image 302 and normalizes the weights to generate a heatmap 308 for GradCAM interpretation.
[0059] Based on the feature vector F and the weight matrix W, the processor 110 uses the weight matrix W to determine the corresponding weighted feature vector F W as the weighted average (dot product) of the corresponding feature vector F along the m×n dimension to obtain a corresponding weighted feature vector F with a dimension of 1×1×d W . In other words, the attribute-weighted feature vector F W is defined as the weighted average of the d-dimensional partial vectors in F, where each feature vector is weighted by its corresponding pixel attribute value in W. The resulting F W has a shape of 1×1×d and is mathematically expressed as follows:
[0060]
[0061] Finally, it is also helpful to know which features contribute the most to the model's decision. Therefore, based on the feature vector F and the weight matrix W, the processor 110 determines each corresponding top attribute feature vector F by using the weight matrix W to determine the maximum weighted value along the m×n dimension in the corresponding feature vector F T , to obtain a corresponding weighted feature vector F with a dimension of 1×1×d T . In other words, the top attribute space feature vector F T is defined as the feature vector corresponding to the region of the maximum GradCAM pixel attributes. F T can be expressed as: F T = F i*,j* , where i*, j* = argmax i,j W i,j . As will be discussed below, F TFor generating a data slice summary with feature visualization.
[0062] Method 200 continues to identify interpretable data slices in the validation dataset (block 230). In particular, the processor 110 identifies a plurality of data slices from a plurality of validation data samples. Each data slice includes a set of validation data samples from the plurality of validation data samples of the dataset 170, and the machine learning model 12 focuses on or depends on similar features for this set of validation data samples to generate predictions. In at least some embodiments, each data slice is an image data slice having a subset of images, and the machine learning model 12 focuses on or depends on similar image features for this subset of images to generate predictions. In at least some embodiments, the processor 110 identifies the plurality of data slices by clustering a plurality of weighted feature vectors F W The processor 110 forms each corresponding data slice into a set of validation data samples that has a corresponding weighted feature vector F W , and these weighted feature vectors are clustered together based on their common or similar features.
[0063] In some embodiments, the processor 110 performs clustering using K-means. In some embodiments, before performing clustering, the processor 110 reduces the dimensionality of the plurality of weighted feature vectors F W to a 2D space. In particular, the processor 110 maps the plurality of weighted feature vectors F W to a two-dimensional space, for example using the Uniform Manifold Approximation and Projection (UMAP) method. This has the effect of reducing the dimensionality and will provide 2D space information for generating data slice mosaics. Parameters in UMAP, such as n_neighbors and min_dist, are fine-tuned to ensure good separation of data subgroups. In some embodiments, to avoid missing or overlooking rare data slices, the processor 110 applies over-clustering by increasing the value of K until a more consistent sample grouping is achieved, as indicated by consistent model behavior within each data slice.
[0064] It should be understood that the resulting data slices enable the system to group validation data samples corresponding to the consistent attributes of the machine learning model 12 together. This is a unique aspect of the slice-based approach utilized by method 200, which is different from traditional clustering analysis. In particular, traditional clustering uses features from the entire image to group data and results in many ambiguous clusters. In contrast, method 200 utilizes an interpretable subset of features with semantics to "slice" the original image data. For example, in the case of a hair color classification model, the data slices will contain images where the model focuses on similar features such as "hair", "mouth", "eyes", "face", or "background". This helps domain or machine learning experts easily diagnose faults in each data slice, who can examine the images and attributes within each data slice. However, as will be understood from the following description, method 200 further reduces human effort by at least partially automating the investigation of individual validation data samples.
[0065] Method 200 continues to generate and display a visualization of at least one data slice (block 240). In particular, the processor 110 generates and displays on the display screen 130 a visualization of at least one data slice from a plurality of data slices. Individually examining the validation data samples in the data slices can be a time-consuming task. Therefore, a more efficient method is needed to provide an overview of the content of each data slice to the user. To address this issue, in at least some embodiments, the visualization takes the form of a data slice mosaic that utilizes feature visualization to generate a visual summary of the data slices in a mosaic representation. In particular, the data slice mosaic includes a plurality of feature visualizations corresponding to the individual data slices, which are arranged as a mosaic. The position of each tile in the data slice mosaic represents the similarity of the features that the machine learning model 12 focuses on or depends on in each data slice compared to other data slices. Each feature visualization is a visual representation of the features that the machine learning model 12 focuses on or depends on, which are used to generate an output for the validation data samples regarding the corresponding data slice.
[0066] The processor 110 operates the display screen 130 to display various graphical user interfaces, which include the data slice mosaic and other tools for exploring and analyzing the data slices. These graphical user interfaces advantageously enable the user to view the aggregated visual patterns of each data slice, validate insights with GradCAM visual explanations, and annotate the revealed issues such as spurious correlations or incorrect labels.
[0067] Figure 5A and Figure 5BExemplary graphical user interfaces 400A, 400B for visualizing analysis data slices are shown. The graphical user interfaces 400A, 400B are shown in the context of an exemplary visual model that has been trained on the CelebA image dataset and is configured to receive an image and output a hair color classification label (e.g., non-gray hair, gray hair) for at least one person represented in the image. However, it should be understood that a graphically similar user interface may be provided regardless of the form or function of the machine learning model 12 to which the method is applied.
[0068] The graphical user interfaces 400A and 400B include a system menu 410 through which a user can select a dataset and a model to be analyzed and verified, and select various visualization layouts and coloring options. The user can choose between two visualization layouts: a combined view or a confusion matrix view, which allows them to obtain an overview of the data slices or examine them in more detail by breaking them down into different error types. The user can also select a coloring matrix that best suits their needs from slice name, slice accuracy, slice confidence, and falsehood probability. The user can also enable or disable scatter plots, feature visualization, and contour visibility.
[0069] The graphical user interface 400A includes a slice table 420 through which a user can navigate and select different data slices for visualization. The slice table 420 enables the user to easily sort the data slices according to various metrics, including accuracy, confidence, and falsehood probability. The processor 110 calculates accuracy and confidence metrics based on the model output regarding the data slices. The processor 110 uses a label propagation method based on user annotations to generate a falsehood probability metric, which will be described in further detail below. The user can sort the data slices according to the selected metric or click on a specific table cell to investigate information about the corresponding slice in other views (e.g., feature visualization or pixel attribute heatmap), which will reconcile the numerical metrics with the qualitative model behavior, thereby enabling more interpretable model verification. The slice table 420 enables the user to identify slice patterns of interest, such as core features with low accuracy, which may imply mislabeled validation data samples in the slice. Alternatively, false slices with high accuracy may indicate that false features can distinguish model predictions and, therefore, the model is not robust to these correlations.
[0070] The graphical user interfaces 400A and 400B include data slice mosaics 430A - 430E. The data slice mosaics 430A - 430E include a plurality of feature visualizations displayed in tiles arranged together to form a mosaic. The data slice mosaics 430A - 430E depict the main visual patterns of each data slice from the perspective of the model. The data slice mosaics 430A - 430E can be displayed in a combined form ( Figure 5AThe data slice mosaics in 430A), or are shown in the form of a confusion matrix that shows separate data slice mosaics with different data subsets segmented by the confusion matrix of the model Figure 5B The data slice mosaics 430B - 430E in show true positives, true negatives, false positives, and false negatives). Data slice mosaic 430B shows the true negative rate of the machine learning model 12 for the data samples in each corresponding data slice. Data slice mosaic 430C shows the false positive rate of the machine learning model 12 for the data samples in each corresponding data slice. Data slice mosaic 430D shows the false negative rate of the machine learning model 12 for the data samples in each corresponding data slice. Finally, data slice mosaic 430E shows the true positive rate of the machine learning model 12 for the data samples in each corresponding data slice. If there is no image in a particular data slice, the corresponding mosaic tile boundary is displayed, colored with the selected metric, to provide a visual context to the user. Additionally, data slice mosaics 430A - 430E also use the color of the mosaic tile boundaries to visualize user - specified metrics such as accuracy, confidence, and falsehood probability, which provides valuable guidance to the user for detecting troublesome slices. The user can easily annotate the slices by double - clicking on the mosaic tile and selecting / entering their annotations.
[0071] Finally, the graphical user interfaces 400A and 400B include a slice detail view 440. The slice detail view 440 displays an image sample of the selected data slice, presenting one or both of the original image and its pixel attribute (GradCAM) heatmap according to the user's selection. In this way, the user can see further details of each slice. Other slice metrics are also displayed in this view, such as slice size and data distribution based on the confusion matrix.
[0072] In at least some embodiments, the slice detail view 440 includes a pixel attribute (GradCAM) heatmap overlaid on the respective validation images. Specifically, the processor 110 generates a pixel attribute heatmap of the respective validation image by mapping the weight matrix W of the corresponding validation image to the pixels of the corresponding validation image. The processor 110 displays the pixel attribute heatmap overlaid on the original validation image in the graphical user interface. In this way, the user can easily understand which features of the image the machine learning model relies on to classify the corresponding validation image.
[0073] Figure 6 summarizes the workflow for generating data slice mosaics for a vision model. As discussed above, the processor 110 uses GradCAM and the machine learning model 12 to compute (block 520) multiple weighted feature vectors F for all the images in the validation dataset 510 W and multiple top - attribute feature vectors FT Next, as discussed above, the processor 110 uses, for example, UMAP to transform the plurality of weighted feature vectors F W Mapping (block 530) to a two-dimensional space. Next, processor 110 performs an eigenvector F on the mapped feature vector F by, for example, using K-means. W Clustering is performed to identify (block 540) a plurality of data slices.
[0074] Next, the processor 110 calculates (block 550) a mapping feature vector F for each of the plurality of data slices. W In particular, to determine the boundaries of the mosaic tiles for each data slice, the processor 110 computes the convex hull of each slice in 2D space. The convex hull is the feature vector F that best maps the cluster / data slice. W Mapping feature vector F to each other cluster / data slice W Separating boundaries. Since the data slices have been clustered in 2D space, the resulting convex hulls have little overlap with each other. This will produce a layout of mosaic tiles that will allow each data slice summary (feature visualization) to be positioned without overlap in the data slice mosaic.
[0075] Next, the processor 110 generates (block 560) a plurality of feature visualizations for the plurality of data slices. Feature visualization is an XAI technique that uses optimization to create images that produce desired responses in specific neurons, channels, or layers of a neural network. Each feature visualization is a visual representation of a feature that the machine learning model 12 is interested in or relies on to generate an output for a validation data sample for the corresponding data slice.
[0076] The processor 110 determines each corresponding feature visualization using feature inversion. Feature inversion is a type of feature visualization and is particularly useful for understanding how a model processes visual information by generating an image that best matches a given representation. By using backpropagation to optimize random values to achieve the same activations as the target image, feature inversion produces outputs that reveal how the network perceives the input image. To achieve this, feature inversion first runs the target image through the network and records the neuron activations at the desired layer. Feature inversion then initializes a new image with random values and uses backpropagation to optimize it to match the target activations. Formally, feature inversion is defined as W×H×C →R d and target feature activation φ(x) = φ 0 In the case of * .
[0077]
[0078] where (l(φ(x), φ 0 ) + λR(x)) is the loss function, which captures the difference between φ(x) and φ 0 , and R(x) is the regularization term.
[0079] In this way, the processor 110 generates a representative image using feature inversion, which outlines the content of multiple images in the data slice. The processor 110 determines the corresponding average value of the top attribute feature vector F T of the corresponding data slice, and sets the target φ 0 optimized by feature inversion to the target of the top attribute feature vector F T of the corresponding data slice, that is, Since the resulting image is an image approximating the top attribute feature vector, the resulting visualization will depict the most dominant visual patterns in the data slice used by the model for prediction.
[0080] In at least one embodiment, feature inversion is performed under the following additional constraint: the generated feature visualization of each data slice has a shape corresponding to the convex hull of the corresponding data slice. In this way, when overlapping the feature visualizations within the mosaic tile boundaries, cropping the feature visualizations will not lose any information. To this end, the processor 110 performs feature inversion with additional constraints. After each iteration, the processor 110 sets the pixels of the resulting image outside the convex hull boundary to 0, forcing the optimization process to focus on the pixels within the convex hull boundary.
[0081] Finally, once the feature visualization is rendered, the processor 110 generates (block 570) a data slice mosaic by arranging the feature visualizations as a mosaic based on the clustering feature vector F W of each data slice and / or the convex hull. Specifically, in the data slice mosaic, each feature visualization is overlapped on the corresponding convex hull to form a mosaic tile. This visualization allows the user to explore the relationships between data slices and identify the key visual patterns that distinguish them. Specifically, each mosaic tile has a feature visualization that helps the user visualize the features of each data slice mainly used by the machine learning model 12 for prediction. In addition, each mosaic tile is arranged at the position representing the features of each data slice in the 2D space of the data slice mosaic, such that based on the proximity of the mosaic tiles to each other, the similarity between data slices can be easily understood.
[0082] As discussed above, data slice mosaics can be displayed in the confusion matrix view. In particular, in the confusion matrix view, the feature visualizations and mosaic tiles are arranged as four different mosaics, each of the four different mosaics representing a different one of (i) the true positive rate, (ii) the true negative rate, (iii) the false positive rate, and (iv) the false negative rate of the machine learning model 12 with respect to each data slice. In particular, if a data slice has samples with corresponding confusion matrix classes and / or error types, the feature visualization is displayed along with the feature visualization in the data slice mosaic. Otherwise, if a data slice does not have samples with corresponding confusion matrix classes and / or error types, the feature visualization is not displayed within the mosaic tile, and only the convex hull is displayed.
[0083] Method 200 continues with annotating slice errors (block 250). In particular, based on user input received from the user, the processor 110 annotates at least one data slice from the plurality of data slices. Based on the at least one annotation of the at least one data slice, the processor 110 determines, using label propagation techniques, the corresponding probabilities of the annotation applying to validation data samples in each of the other data slices from the plurality of data slices. Thus, either a manual annotation from the user or a propagated annotation probability is provided to each data slice in the plurality of data slices. By leveraging the annotation probabilities, the amount of human effort required to annotate the validation data set is greatly reduced.
[0084] In at least some embodiments, each annotation includes one of (1) a spurious feature label or (2) a core feature label. As described above, a spurious feature label applied to a data slice indicates that the machine learning model 12 focuses on or relies on spurious features to generate an output for the validation data samples of that data slice. In other words, the machine learning model 12 associates the wrong features with the output label. In the example of a hair color classification model, spurious features may include the background, facial features, or any other features in the image that are not hair features. In contrast, a core feature label applied to a data slice indicates that the core features are focused on or relied on by the machine learning model 12 to generate an output for the validation data samples of that data slice. In other words, the machine learning model 12 associates the correct features with the output label. In the example of a hair color classification model, the core features would include the hair features within the image. In at least some embodiments, some annotations may also include an incorrect label, which indicates that the validation data samples within the data slice include incorrectly labeled validation data samples.
[0085] It should be understood that the graphical user interfaces 400A, 400B, and in particular the data slice mosaics 430A - E, make it easy for the user to identify problems within the data slices, such as spurious correlations. Spurious correlations can exist in any machine learning model, regardless of accuracy, and can lead to significant problems, including but not limited to poor generalization performance in production or AI fairness issues.
[0086] Referring again to Figure 5A and Figure 5B when the user navigates the graphical user interfaces 400A, 400B, the user can click on the mosaic tiles within the data slice mosaics 430A - E to annotate the corresponding data slices with annotations / labels. In response to such selection of a mosaic tile, the processor 110 displays an annotation window 450 within the graphical user interfaces 400A, 400B. The user can interact with the annotation window 450 to apply a spurious feature label or a core feature label. Additionally, the user can apply descriptive labels to further classify the data slices (e.g., "core: grey hair", "bone spur: mouth", or "false label").
[0087] Once at least one data slice has been annotated with a spurious feature label or a core feature label, the processor 110 determines the spuriousness probability for each other data slice that has not been annotated with a spurious feature label or a core feature label. The so-called "spuriousness probability" is a value ranging from 0 to 1 that indicates the probability that the machine learning model 12 focuses on or relies on spurious features to generate an output of a validation data sample for that data slice. Thus, a spuriousness probability close to 0 indicates that the machine learning model 12 is likely to use the core features of the data slice, while a spuriousness probability close to 1 indicates that the machine learning model 12 is likely to use the spurious features of the data slice.
[0088] The processor 110 uses a label propagation method, such as the scikit - learn method, to calculate the spuriousness probability for each unannotated data slice, which automatically generates this probability for the unannotated data slices based on the similarity between the user's annotations and the data slice feature representations. In particular, in some embodiments, the processor 110 determines the feature representation of each data slice as the average of the weighted feature vectors F W of the validation data samples within the data slice. Next, the processor 110 calculates the spuriousness probability for each unannotated data slice based on the similarity (e.g., distance) between the feature representation of the unannotated data slice and the feature representations of the annotated data slices.
[0089] Referring again to Figure 5A and Figure 5B the spuriousness probability is displayed to the user in the graphical user interfaces 400A, 400B. First, the slice table 420 allows the user to sort the data slices by spuriousness probability. Additionally, the outline or convex hull of each mosaic tile within the data slice mosaics 430A - E is color - coded according to the spuriousness probability. In this way, the spuriousness probability makes the data slice exploration process easier because slices with hypothesized spurious correlations are highlighted. This is an important step in helping the user detect and evaluate problematic slices.
[0090] Method 200 continues to retrain the machine learning model to reduce the dependence on spurious features (block 260). Specifically, an effective way to reduce spurious correlations in machine learning model 12 is through model retraining, which can improve the model's robustness to potential biases without changing the architecture. As discussed above, in addition to multiple validation data samples, dataset 170 also includes multiple training data samples for training machine learning model 12. Based on the artificial spuriousness annotations and the calculated spuriousness probabilities, processor 110 modifies one or more of the multiple training data samples in a manner designed to prevent machine learning model 12 from erroneously learning to rely on spurious features to generate outputs regarding those training data samples. Once the multiple training data samples are modified as needed, processor 110 uses the modified multiple training data samples to retrain machine learning model 12.
[0091] To identify which training data samples should be modified, processor 110 identifies which of the multiple data slices have spurious feature labels or have a spuriousness probability exceeding a predetermined threshold, i.e., the data slices corresponding to a significant dependence of machine learning model 12 on spurious features. Next, processor 110 identifies the training data samples corresponding to or similar to the identified data slices of the validation data samples. In one embodiment, processor 110 determines the weighted feature vectors of the multiple training data samples in the same manner as the multiple weighted feature vectors F W used to determine the multiple validation data samples. Next, processor 110 determines to which data slice each training data sample belongs by mapping the respective weighted feature vectors to the clusters of the weighted feature vectors F W of the data slices. Once the set of training data samples is identified for modification, processor 110 modifies those training data samples in a manner designed to prevent machine learning model 12 from erroneously learning to rely on spurious features to generate outputs regarding those training data samples. In at least some embodiments, in addition to modifying the training data to reduce the dependence on spurious features, if some training data has been identified as having incorrect labels, processor 110 also relabels this training data.
[0092] In the case of a vision model trained using training images, in some embodiments, the Core Risk Minimization (CoRM) method is utilized to reduce the model's dependence on spurious features. CoRM uses random Gaussian noise to corrupt non-core image regions and retrains the model using the data corrupted by the noise, which has been shown to be effective in reducing the model's dependence on spurious features. Specifically, processor 110 applies noise to the portions of the identified training images corresponding to spurious features and uses the training images corrupted by the noise to retrain machine learning model 12.
[0093] Figure 7 An exemplary workflow 600 for training images in which noise disrupts spurious features is outlined. For a training image corresponding to a problematic data slice, pixel attributes (i.e., GradCAM masks or heatmaps 620) are used to highlight spurious regions, and these masks are used to add random Gaussian noise to the spurious regions. Specifically, the processor 110 determines a weight matrix for each training image 610 to be modified in a manner similar to that discussed above with respect to determining the weight matrix W. The processor 110 uses the GradCAM method to map the weight matrix to the pixels of the training image to determine the GradCAM mask m. Finally, the processor 110 applies Gaussian noise to the training image using the mask m to provide a noise-disrupted training image 630. For a single image, this process can be represented as x' = x + m ⊙ z, where x is the input image, m is the GradCAM mask, and z is the generated Gaussian noise matrix. All three variables are the same size as the input image, and ⊙ represents the Hadammard product. Figure 7 Some examples of this operation are shown, with the noise exaggerated for demonstration purposes.
[0094] After replacing the original training data samples with modified training data samples (e.g., noise-disrupted images), the processor 110 retrains the machine learning model 12, which should result in a trained machine learning model 12 that has a reduced dependence on spurious features and provides better generalization performance.
[0095] Illustrative Case Study
[0096] To illustrate how the insights of method 200 can be used to improve machine learning models, two case studies are discussed. These case studies utilize visual models trained on publicly available visual datasets to benchmark and evaluate the capabilities of method 200. The main purpose of these case studies is to demonstrate how method 200 enables machine learning experts and practitioners to detect, evaluate, and interpret potential problems in visual models.
[0097] Figure 8 A common usage pattern 700 of the system and method is shown. In particular, a user can start an analysis using the slice table of the graphical user interface for a rank-driven evaluation, where data slices are ranked by a score (e.g., model accuracy or spuriousness), or start an analysis using the data slice mosaic of the graphical user interface for a visually driven evaluation. Once a data slice of interest is identified, individual validation data samples can be explored in the slice details view of the graphical user interface. Finally, the user can annotate the data slice and continue their analysis.
[0098] In the first case study, Method 200 was applied to a hair color classification model to find edge cases in a dataset. This case study is a previous illustrative example of the Figure 1 , Figure 4 , Figure 5A , Figure 5B and Figure 6 used. This case study involves the large-scale CelebA (CelebA) dataset with 202,599 facial images. The label for each image is one of the classifications {non-gray hair, gray hair}, called labels {0, 1} respectively. Through an 8:1:1 training, validation, and test split, transfer learning was used to train a ResNet50 binary image classifier. After iteratively fine-tuning the hyperparameters, a trained model with a classification accuracy of 98.03% was obtained. The experts used Method 200 to interpret and diagnose the performance of this hair color classifier. For the UMAP algorithm, it was set to n_neighbors = 5, min_dist = 0.01, and n_components = 2, and for K-means clustering, n_clusters = 50.
[0099] Does the model behave correctly on slices with good performance? The expected behavior of a basic hair color classification model is to capture hair features. Referring again to Figure 5A , upon first looking at the data slice mosaic 430C, the experts noticed that slice_22 and slice_5 on the left were separated from the other slices on the data slice mosaic. Their feature visualizations were suggestive of a gray hair pattern, which means the model was using the correct features, i.e., the core features. By examining the corresponding pixel attributes using the slice detail view 440, they confirmed the correctness of this insight and annotated the two data slices with the description "core: gray hair" as core features. When the experts saved the annotations, Method 200 automatically propagated the annotations and provided the falsity probability for each slice in the graphical user interface 400A. Under this guidance, the experts observed that many slices located on the right side of the data slice mosaic 430A had a higher falsity probability. Moreover, their feature visualizations did not suggest a hair pattern. Instead, many feature visualizations seemed to include other facial features such as eyes and mouths, leading to a valid suspicion that the model was not performing correctly on these slices. In addition, the experts also noticed that the model had correct predictions for these slices (prediction accuracy of 100%). This gave the experts reason to worry that the model was largely biased by spurious features.
[0100] Figure 9Shows some exemplary spuriousness issues 800 detected during the operation of the hair color classifier model. In particular, through investigation, experts found that the model incorrectly utilized the mouth and eyes to predict the hair color of several of the best-performing slices, such as slice_0, slice_1, and slice_2. As can be seen from the heatmap 800, the model focuses on spurious correlations related to the image background, mouth, face, and eyes. Additionally, some image samples were identified with incorrect labels.
[0101] Why do some slices perform poorly? The experts were also eager to investigate the underperforming data slices. Referring again to Figure 5A , by sorting the slice table 420 in ascending order of accuracy, they noticed that the model achieved only 72.41% accuracy on slice_44 and found that the feature visualization of this slice only showed meaningless colored patterns. After examining the pixel attributes, they found that the model looked at the image background for prediction, which can be seen in the heatmap 800 of Figure 9 . This spurious correlation problem was prominent, and they annotated this slice as "spurious" with the description "spurious: background". Similar issues also occurred in its adjacent slices, and method 200 automatically assigned higher spuriousness probabilities to them.
[0102] What potential factors contribute to unexpected behavior ? Machine learning experts were interested in understanding why this unexpected model behavior occurred. Referring again to Figure 5A , by sorting the slice table 420 in descending order of spuriousness, they obtained a list of slices with high spuriousness probabilities. They switched the data slice mosaics into the form of a confusion matrix, as shown in Figure 5B , to further study the details. By clicking on the name of "slice_35", they highlighted this slice on the four data slice mosaic subviews 430B - E and examined the provided explanations, where they noticed that the images in the "false negative" group located in this slice had incorrect labels - they should have been labeled "non - grey hair" instead of "grey hair". They marked this problem as "incorrect label". By investigating its adjacent slices, they discovered and marked another slice with an incorrect label, slice_2_FN, through several clicks.
[0103] In summary, through the above procedures in this case study, it was demonstrated how method 200 supports users in revealing and explaining potential model problems through visual summaries and helpful guidance. Based on the annotations from machine learning experts, method 200 adopts the CoRM framework to mitigate the detected errors, which will be evaluated below.
[0104] In the second case study, the method 200 was applied to a bird category classification model to detect bias. In particular, to study whether the method 200 can help machine learning experts and practitioners find potential biases and discriminative power of the model, this case study was designed with a biased dataset called water birds, which was constructed by cropping birds from the photos in the Caltech-UCSD Birds-200-2011 (CUB) dataset and transferring them to the background from the location dataset. For each image, the label belongs to one of {water birds, land birds}, and the image background belongs to one of {water background, land background}. The training set was skewed by placing 95% of the water birds (land birds) on the water (land) background and the remaining 5% on the land (water) background. The training, validation, and test sets include 4,795, 1,199, and 5,794 images, respectively. After training and fine-tuning the hyperparameters, the water bird / land bird classification model achieved a classification accuracy of 85.74%. In data slice finding, n_neighbors = 20, min_dist = 0.05, and n_components = 2 were set for the UMAP algorithm, and n_clusters = 50 was set for K-means.
[0105] In this study, the machine learning expert realized that the model was likely to be biased by the background, i.e., classifying water birds using the water background, and the same was true for land birds. However, this prior knowledge was difficult to establish in real-world applications due to the lack of additional well-labeled metadata. Therefore, the expert assumed that this information was unknown and wanted to verify whether the method 200 could highlight potential model biases by only using the original images and the trained model.
[0106] Figure 10A and Figure 10B show further exemplary graphical user interfaces 900A, 900B for visualizing and analyzing data slices. The graphical user interfaces 900A, 900B are essentially similar to Figure 5A and Figure 5B the graphical user interfaces 400A, 400B shown in
[0107] Does the model exhibit bias? To answer this key question, experts began by searching for and investigating underperforming slices. From slice table 920, they sorted the listed slices in ascending order of accuracy and selected the worst-performing slice_30. Figure 11 Some exemplary spuriousness issues 1000 detected during the operation of the bird classifier model are shown. In particular, the coordinated information provided by feature visualization and pixel attributes highlights spurious correlation issues, where the model incorrectly uses water and land backgrounds to classify birds. The experts annotated this slice as using spurious features, water background, and method 200 automatically propagated this annotation. They verified the correctness of the propagation on adjacent slices and also annotated slice_2 as using the same spurious features. Additionally, the experts identified underperforming slice_42, which did not cluster with other slices and was given a high spuriousness probability. They studied this slice and verified that the model used spurious features, land background, to predict the bird classification in this slice. Refer Figure 10B , through the data slice mosaics 930B-E in the confusion matrix view, the experts found that this spurious correlation led to many false negatives, where the model used the land background to incorrectly predict many "water birds" as "land birds".
[0108] Is the detected bias prevalent in all slices? Why or why not? Machine learning experts are interested in determining whether the detected bias is prevalent throughout the dataset. Refer Figure 10A , by analyzing slices that are far from the annotated slice in data slice mosaic 930A and are assigned a low spuriousness probability, the experts found that the farthest neighbors, namely slice_24 and slice_14, corresponded to core features. This indicates that the model can correctly capture the bird regions in these slices, which raises a subsequent "why" question. To understand under what circumstances the model fails, the experts used detail view 940 to browse the original images from slice_42 (spurious features) and slice_14 (core features) separately. They found that slice_42 had a very similar land background and very different birds, while on the other hand, the birds in slice_14 (core features) had a very consistent appearance. This finding explains why this biased model can still capture the core features from slice_14 - the greater the similarity of the core features in the representation space, the stronger the robustness to spurious correlations. This insight helps improve the robustness of the model and has been further studied by machine learning experts. In summary, method 200 enables machine learning experts to verify the existence of model bias and extract slices corresponding to different biases.
[0109] Evaluation
[0110] The method 200 was also quantitatively and qualitatively evaluated through two case studies, demonstrating its effectiveness in model validation for visual tasks. The method 200 can help researchers and practitioners better understand and mitigate edge cases in visual applications, ultimately leading to more reliable and accurate machine learning models. Both quantitative and qualitative evaluations were conducted to verify whether the method 200 can indeed leverage human insights to enhance the performance of visual models while reducing their reliance on spurious features.
[0111] The quantitative evaluation involves four metrics, including: net accuracy, core accuracy, spurious accuracy, and relative core sensitivity. Net accuracy is the model accuracy calculated based on the original dataset, where a larger value indicates better overall accuracy. Core accuracy, acc (C) , is the model accuracy calculated when the spurious regions are masked by Gaussian noise, where a larger value indicates a greater reliance of the model on the core regions. Spurious accuracy, acc (S) , is the model accuracy calculated when the core regions are masked by Gaussian noise, where a larger value indicates that the model relies more on spurious regions. RCS is a metric that quantifies the model's dependence on core features while controlling the overall noise robustness. RCS is defined as the ratio of the absolute difference between core accuracy and spurious accuracy to the total possible difference between core accuracy and spurious accuracy for any model, and is expressed as where The RCS ranges from 0 to 1, where a higher value indicates better model performance.
[0112] In the two case studies discussed above, machine learning experts annotated five spurious slices {slice_0, slice_1, slice_2, slice_44, slice_26} in the CelebA dataset and six spurious slices {slice_30, slice_2, slice_22, slice_12, slice_42, slice_9} in the waterbird dataset, respectively. For each case study, the method 200 automatically runs the label propagation algorithm and exports both the user's annotation records and the propagated spuriousness probabilities for further investigation.
[0113] To thoroughly evaluate the method 200, three models were evaluated for each case. The model labeled "baseline" is the original trained model obtained at the start of each case study. After adding noise to the "spurious" slices according to the results derived from the method 200, the model labeled "AS" was retrained using the CoRM method. Specifically, "annotated" indicates that only the spurious slices annotated by the user were corrupted by noise, while "propagated" indicates that the propagated spuriousness was used to identify the spurious slices that would be modified by noise.
[0114] Figure 12Table 1100 is shown, which outlines the quantitative evaluation of the overall performance of the hair color classification model trained on the original CelebA dataset and the modified CelebA dataset improved using Method 200, and the dependence on spurious features. Performance is evaluated using the validation set.
[0115] Figure 13 Table 1200 is shown, which outlines the quantitative evaluation of the overall performance of the bird category classification model trained on the original waterbird dataset and the modified waterbird dataset improved using Method 200, and the dependence on spurious features. Performance is evaluated using the validation set.
[0116] As can be seen, Method 200 can significantly improve the overall performance of the visual model and reduce spurious correlations. In addition, the label propagation of Method 200 significantly reduces the manual effort through the automated annotation process and achieves the best performance in this quantitative evaluation. Overall, the results demonstrate that Method 200 is effective in mitigating spurious correlations in machine learning models, and the label propagation algorithm is a valuable tool for the automated annotation process.
[0117] To qualitatively evaluate the results, GradCAM is used to compare the properties of the models. Figure 14 The qualitative evaluation using the pixel properties from GradCAM is shown. In summary 1300A, a visual comparison of the original classification model and the improved model is made for hair color classification on the CelebA dataset. In summary 1300B, a visual comparison of the original classification model and the improved model is made for bird category classification on the waterbird dataset. For each subplot, the first row refers to the original model, and the second row refers to the improved model. As can be seen, the retrained and improved model suppresses spurious correlations by making predictions using the correct features.
[0118] The hair color classification model initially had major spurious correlations in {Slice_44, Slice_0, and Slice_1}, where the model used the image background, mouth, or eyes respectively to predict hair color. With the help of Method 200, as Figure 14 shown in summary 1300A of, the model was successfully retrained to focus on the correct hair region. As for the bird category classification model, it initially focused on spurious features such as the water / land background to decide whether there was a water / land bird in the input image {Slice_30, Slice_2, and Slice_42}. As Figure 14 shown in summary 1300B of, Method 200 alleviates these problems by helping the model focus on the core bird features.
[0119] Embodiments within the scope of the present disclosure may also include a non-transitory computer-readable storage medium or machine-readable medium for carrying or having stored thereon computer-executable instructions (also referred to as program instructions) or data structures. Such a non-transitory computer-readable storage medium or machine-readable medium can be any available medium accessible by a general-purpose or special-purpose computer. By way of example and not limitation, such a non-transitory computer-readable storage medium or machine-readable medium can include RAM, ROM, EEPROM, CD-ROM or other optical disk storage, magnetic disk storage or other magnetic storage devices, or any other medium that can be used to carry or store the desired program code means in the form of computer-executable instructions or data structures. Combinations of the above should also be included within the scope of the non-transitory computer-readable storage medium or machine-readable medium.
[0120] For example, computer-executable instructions include instructions and data that cause a general-purpose computer, special-purpose computer, or special-purpose processing device to perform a certain function or a set of functions. Computer-executable instructions also include program modules executed by a computer in a stand-alone or networked environment. Generally, program modules include routines, programs, objects, components, and data structures, etc., which perform specific tasks or implement specific abstract data types. Computer-executable instructions, associated data structures, and program modules represent examples of program code means for performing the method steps disclosed herein. A particular sequence of such executable instructions or associated data structures represents an example of the corresponding actions for implementing the functions described in these steps.
[0121] Although the present disclosure has been described in detail in the drawings and the foregoing description, it should be regarded as illustrative rather than restrictive in nature. It should be understood that only the preferred embodiments are presented, and all changes, modifications, and further applications within the spirit of the present disclosure are expected to be protected.
Claims
1. A method for validating a visual model, the visual model being configured to receive an input image and provide at least one output label, the method comprising: storing a plurality of verification images in a memory; determining, with a processor, a plurality of first feature vectors using the vision model, the plurality of first feature vectors comprising a respective first feature vector for each respective verification image in the plurality of verification images; identifying, using a processor, a plurality of groups of verification images from the plurality of verification images based on the plurality of first feature vectors; as well as A visualization of at least one set of verification images from the plurality of sets of verification images is displayed on a display screen.
2. The method of claim 1, wherein the visual model is a neural network having a plurality of layers, and determining the corresponding first feature vector for each corresponding verification image further comprises: Input the corresponding verification image into the visual model; extracting a corresponding intermediate output of at least one intermediate layer among the plurality of layers of the vision model; as well as A respective first eigenvector is determined based on the respective intermediate output.
3. The method of claim 2, wherein determining the corresponding first feature vector of each corresponding verification image further comprises: Determining a weight matrix based on a model gradient of the visual model at at least one intermediate layer; as well as The corresponding first eigenvector is determined by determining a weighted average of the corresponding intermediate output along at least one dimension using a weight matrix.
4. The method according to claim 3, further comprising: Generate a pixel attribute heat map of the corresponding verification image in the plurality of verification images by mapping the weight matrix of the corresponding verification image to the pixel of the corresponding verification image; as well as Show a heat map of pixel attributes overlaid on the corresponding validation image on the display.
5. The method according to claim 1, identifying multiple sets of verification images further comprises: Clustering the multiple first eigenvectors; as well as Each respective verification image group is formed from those verification images corresponding to respective first feature vector groups clustered together among the plurality of first feature vectors.
6. The method according to claim 1, further comprising: A plurality of second feature vectors including a respective second feature vector for each respective verification image in the plurality of verification images is determined, with the processor, by determining a maximum weighted value along at least one dimension in the respective intermediate outputs using a weight matrix.
7. The method of claim 6, wherein displaying the visualization of at least one set of verification images further comprises: A respective visualization of the respective set of verification images is generated based on a respective set of second feature vectors from the plurality of second feature vectors corresponding to the respective set of verification images using the visual model.
8. The method of claim 7, generating a corresponding visualization of the corresponding verification image set further comprising: determining a corresponding average second feature vector for the corresponding verification image group by averaging a group of second feature vectors corresponding to the corresponding verification image group from the plurality of second feature vectors; as well as Based on the average second eigenvector, a representative image is generated as the corresponding visualization, wherein the representative image is generated via a feature inversion technique using a visual model.
9. The method of claim 1, displaying a visualization of at least one set of verification images further comprising: Mapping a plurality of first eigenvectors onto a two-dimensional plane; generating a plurality of visualizations, each respective visualization in the plurality of visualizations representing a respective validation image group in the plurality of sets of validation images; as well as A plurality of visualizations arranged as a mosaic are displayed on a display screen, wherein each respective visualization is arranged relative to the other visualizations according to a mapping of the plurality of first eigenvectors onto the two-dimensional plane.
10. The method of claim 9, wherein each respective one of the plurality of visualizations has a shape and position in the mosaic corresponding to a convex hull of a cluster of respective first eigenvectors of the plurality of first eigenvectors on a two-dimensional plane, the respective first eigenvectors corresponding to the respective verification image group represented by the respective visualization.
11. The method of claim 9, displaying a plurality of visualizations arranged as a mosaic further comprising: A plurality of visualizations arranged as four different mosaics are displayed on a display screen, each of the four different mosaics representing a different one of (i) true positive rate, (ii) true negative rate, (iii) false positive rate, and (iv) false negative rate of the visual model with respect to each of a plurality of sets of validation images.
12. The method according to claim 1, further comprising: annotating at least one corresponding set of verification images from the plurality of sets of verification images with at least one annotation based on user input; as well as Based on the at least one annotation, a respective probability that the at least one annotation applies to an image in each respective set of validation images in the plurality of sets of validation images is determined.
13. The method of claim 12, wherein at least one annotation indicates that the vision model erroneously used a spurious feature to determine at least one output label for an image in at least one corresponding validation image set.
14. The method according to claim 12, further comprising: storing in a memory a plurality of training images for training a vision model; modifying, with a processor, a training image of the plurality of training images based on the at least one annotation and a corresponding probability that the at least one annotation applies to each of the plurality of sets of validation images; as well as The processor is used to retrain the vision model using the modified plurality of training images.
15. The method of claim 14, wherein modifying a training image in the plurality of training images further comprises: identifying a group of verification images in the plurality of verification images to which a corresponding probability of at least one annotation being applied exceeds a predetermined threshold; According to the identified verification image group, identifying a training image to be modified among the plurality of training images; as well as Noise is applied to the identified portion of the plurality of training images.
16. The method according to claim 15, identifying a training image to be modified from among the plurality of training images further comprising: For each corresponding training image in the plurality of training images, determining a corresponding feature vector; as well as For each respective one of the plurality of training images, responsive to determining that the respective feature vector maps to a cluster of first feature vectors of the plurality of first feature vectors corresponding to an identified validation image group of the plurality of validation images, identifying the respective training image to be modified.
17. The method of claim 16, wherein the vision model is a neural network having a plurality of layers, and applying noise to the identified portion of the training image further comprises: For each corresponding training image to be modified, determining a weight matrix based on a model gradient of the vision model at at least one intermediate layer; For each corresponding training image to be modified, mapping the weight matrix to pixels of the corresponding training image; as well as For each corresponding training image to be modified, noise is applied to the pixels of the corresponding training image according to the mapping of the weight matrix.
18. The method according to claim 14, further comprising: At least one training image of the plurality of training images is relabeled.
19. A method for validating a machine learning model, the machine learning model being configured to receive an input data sample and provide at least one output label, the method comprising: storing a plurality of verification data samples in a memory; determining, with a processor, a plurality of first feature vectors using a machine learning model, the plurality of first feature vectors comprising a respective first feature vector for each respective validation data sample in a plurality of validation data samples; identifying, using a processor, a plurality of groups of validation data samples from the plurality of validation data samples based on the plurality of first feature vectors; and A visualization of at least one set of validation data samples from the plurality of sets of validation data samples is displayed on a display screen.
20. A system for validating a machine learning model, the machine learning model being configured to receive an input data sample and provide at least one output label, the system comprising: a display screen configured to display a graphical user interface; at least one memory configured to store program instructions for the machine learning model, a plurality of validation data samples, and a plurality of training data samples; and At least one processor is configured to perform the following operations: determining, using the machine learning model, a plurality of first feature vectors, the plurality of first feature vectors comprising a respective first feature vector for each respective validation data sample in a plurality of validation data samples; Identifying a plurality of groups of validation data samples from the plurality of validation data samples; displaying within a graphical user interface a visualization of at least one set of validation data samples from the plurality of sets of validation data samples; annotating at least one respective set of validation data samples from the plurality of sets of validation data samples with at least one annotation based on the user input; modifying a training data sample from the plurality of training data samples based on the at least one annotation; as well as Retrain the machine learning model using multiple modified training data samples.