Method and device for training user classification model
By using the classification model initialized by the pre-trained model to divide and classify the target domain data without accessing the source domain data, generate pseudo-labels and update the model, the problem of unstable performance of the classification model on the target domain data is solved, and the model adaptation and identification of unknown categories are realized.
Patent Information
- Application Number
- CN202510220527.3
- Authority / Receiving Office
- CN · China
- Patent Type
- Applications(China)
- Current Assignee / Owner
- Filing Date
- 2025-02-26
- Publication Date
- 2025-05-30
AI Technical Summary
In practical applications, the classification model faces significant differences in the number of categories and probability distribution of data classification between the training data and the test data in actual application scenarios, resulting in the unstable performance of the model on the target domain data. Especially when the source domain data is inaccessible, the traditional domain adaptation method cannot be applied.
By obtaining two classification models initialized based on the pre-trained model, class clustering and classification prediction are performed on the tagless user set, class cluster pseudo-labels and classification pseudo-labels are generated, users of unknown categories are determined, and users of unknown categories are updated by using these users to realize the adaptation of the model to the target domain data.
Without accessing source domain data, the model can be adapted to the target domain data, achieved accurate classification, and identified unknown categories in the target domain data, improving the flexibility and generalization capabilities of the model.
Smart Images

Figure CN120067842A_ABST
Abstract
Description
Technical Field
[0001] One or more embodiments of this specification relate to the field of machine learning technology, and in particular, to a method and device for training a user classification model. Background Art
[0002] The rapid development of machine learning has promoted the wide application of various machine learning models in diverse business scenarios. For example, in the financial field, in order to provide personalized financial products that meet the personal financial preferences of a large number of customers, a well-trained machine learning model is often applied to the classification of customer portraits (user samples). Such models can identify user sample features and achieve accurate user classification.
[0003] However, with the rapid expansion of business, user samples are constantly changing. For example, the addition of new user samples, the data update of historical user samples, and so on. In actual applications, classification models usually face a challenge, that is, there are significant differences in the number of data classification categories and probability distributions between the training data (also called source domain data) and the test data (also called target domain data) in the actual application scenario.
[0004] For this reason, the Domain Adaptation (DA) method has emerged, which aims to solve the inconsistency in feature distributions between the source domain data used in the training phase of the model and the target domain data in actual applications. Through the DA method, the knowledge learned by the model on the source domain data can be transferred to a new application scenario, ensuring that the model can also maintain stable performance on the target domain data.
[0005] Nevertheless, the DA methods in many related technologies rely on access to the source domain data, and this prerequisite is often difficult to meet in actual application scenarios. For example, when the source domain data involves sensitive privacy information, or there are data access restrictions in laws and regulations, etc., the source domain data is not allowed to be accessed. At this time, traditional DA methods are not applicable.
[0006] Therefore, it is hoped that there can be a technical solution that can train the model under the condition of not accessing the source domain data of the pre-trained model, make it adapt to the target domain, perform accurate classification on the target domain data, and be able to effectively identify unknown categories in the target domain data. Summary of the Invention
[0007] One or more embodiments of this specification describe a method and device for training a user classification model, which can train the classification model without accessing the source domain data, make it adapt to the target domain, perform accurate classification on the target domain data, and at the same time can identify unknown categories existing in the target domain data, improving the flexibility and generalization ability of the classification model in actual applications.
[0008] According to a first aspect, a method for training a user classification model is provided, including:
[0009] Obtain a first classification model and a second classification model respectively initialized based on a pre-trained model, where the pre-trained model is trained based on a data set of M user groups.
[0010] Perform cluster partitioning on an unlabeled user set, and use the cluster partitioning result to perform a first update on the first classification model.
[0011] Use the first classification model after the first update to perform cluster prediction on each user in the first subset of the user set to obtain cluster pseudo-labels; and use the second classification model to perform classification prediction of the M user groups on each user to obtain classification pseudo-labels.
[0012] Determine a first user with an unknown category from the first subset according to the cluster pseudo-labels and classification pseudo-labels of each user.
[0013] Use the first / second classification model to respectively predict the augmented samples of the first user to obtain first / second prediction results; aim at the first prediction result being close to the cluster pseudo-label of the first user to update the first classification model; aim at maximizing the sum of the prediction probabilities corresponding to the M user groups in the second prediction result to update the second classification model.
[0014] According to an embodiment, the data set comes from a first business domain, and the user set comes from a second business domain.
[0015] According to an embodiment, before the first update, the method further includes:
[0016] Perform data augmentation on the user set to obtain an augmented sample set; based on the user set and the augmented sample set, use a contrastive learning method to train the first classification model.
[0017] According to an embodiment, the performing cluster partitioning on the unlabeled user set includes:
[0018] For the user set, construct a K-nearest neighbor graph, which includes nodes representing users and connection edge weights representing the similarity between the nodes.
[0019] Process the K-nearest neighbor graph based on a community discovery algorithm to obtain several communities.
[0020] Determine a first K value according to the first cluster evaluation of the communities obtained under different K values.
[0021] Using several communities included in the K-nearest neighbor graph constructed with the first K value as the clustering result.
[0022] According to an implementation of the above embodiment, the community discovery algorithm is the Louvain algorithm, and the first cluster evaluation is the within-cluster sum of squares evaluation.
[0023] According to an embodiment, the clustering result indicates the target clusters to which each user belongs; the first update includes:
[0024] Performing cluster prediction on each user in the second subset of the user set to obtain the prediction vectors of each user, including the prediction probabilities of the user belonging to each cluster.
[0025] Comparing the prediction probability corresponding to the target cluster in the prediction vector with a preset first threshold, and dividing the users in the second subset into member users and non-member users according to the comparison result.
[0026] Taking the maximization of the sum of the first probability and the second probability as the goal, updating the first classification model; the first probability is positively correlated with the prediction probability of the member users on the corresponding target cluster; the second probability is positively correlated with the sum of the prediction vector similarities between the non-member users and their respective neighbor samples.
[0027] According to an embodiment, obtaining the classification pseudo-labels includes:
[0028] Obtaining the classification prediction results of each user output by the second classification model, and classifying each user into the group set of the user group corresponding to the highest prediction probability in the classification prediction results.
[0029] Based on the group set, determining the belonging threshold and the exclusion threshold of the corresponding user group.
[0030] For any user in any group set, if its highest prediction probability is higher than the belonging threshold, then mark the classification pseudo-label of the user as the user group corresponding to the group set; if its highest prediction probability is lower than the exclusion threshold, then mark the classification pseudo-label of the user as an unknown category.
[0031] According to an implementation of the above embodiment, based on the group set, determining the belonging threshold and the exclusion threshold of the corresponding user group includes:
[0032] For any group set, based on the maximum likelihood estimation algorithm, fitting the belonging probability distribution and the exclusion probability distribution.
[0033] According to the expected probabilities of the belonging / exclusion probability distributions, respectively determining the belonging threshold and the exclusion threshold of the user group corresponding to the group set.
[0034] According to one implementation, the maximum likelihood estimation algorithm is the Expectation-Maximization (EM) algorithm, and the attribution probability distribution and the exclusion probability distribution are beta distributions.
[0035] According to one embodiment, the classification pseudo-labels include an unknown class label and a known class label representing one of the M user groups; determining the first user with an unknown class from the first subset includes:
[0036] Determining the target mapping relationship between the cluster pseudo-labels and the M user groups, such that the number of consensus users having cluster pseudo-labels that conform to the classification pseudo-labels according to the target mapping relationship is maximized.
[0037] Determining the set of users with unknown classes, which includes users with classification pseudo-labels being unknown class labels, and users with classification pseudo-labels being known class labels but not belonging to the consensus users; the first user is selected from the set of users with unknown classes.
[0038] According to one implementation of the above embodiments, the method further includes:
[0039] Using the first / second classification model to respectively predict the augmented samples of any consensus user, obtaining the first / second consensus item prediction results; aiming at the first / second consensus item prediction results being close to the cluster / classification pseudo-labels of the consensus user, respectively updating the first / second classification model.
[0040] According to one implementation of the above embodiments, the method further includes:
[0041] For any user to be determined, obtaining the prediction result vectors of the user to be determined and its respective neighbor samples output by the first classification model, and updating the first classification model with the goal of maximizing the similarity between the prediction result vectors; the users to be determined include other users in the first subset who do not belong to the consensus users nor the set of users with unknown classes.
[0042] Using the second classification model to respectively predict the samples of any user to be determined after strong augmentation and weak augmentation; and updating the second classification model with the goal of maximizing the consistency of the prediction results.
[0043] According to a second aspect, there is provided an apparatus for training a user classification model, including:
[0044] An acquisition module, configured to acquire a first classification model and a second classification model respectively initialized based on a pre-trained model, and the pre-trained model is trained based on a data set of M user groups.
[0045] A first update module, configured to perform cluster partitioning on the set of unlabeled users, and use the cluster partitioning result to perform a first update on the first classification model.
[0046] A prediction module, configured to use the first updated first classification model to perform cluster prediction on each user in the first subset of the user set to obtain cluster pseudo-labels; and use the second classification model to perform classification prediction on each user for the M user groups to obtain classification pseudo-labels.
[0047] A determination module, configured to determine a first user with an unknown category from the first subset according to the cluster pseudo-labels and classification pseudo-labels of each user.
[0048] A second update module, configured to use the first / second classification model to respectively predict the augmented samples of the first user to obtain first / second prediction results; target that the first prediction result is close to the cluster pseudo-label of the first user, and update the first classification model; target that the sum of the prediction probabilities corresponding to the M user groups in the second prediction result is maximized, and update the second classification model.
[0049] According to a third aspect, there is provided a computer program product, including a computer program / instructions, which when executed by a processor, implement the steps of the method described in the first aspect.
[0050] According to a fourth aspect, there is provided a computing device, including a memory and a processor, characterized in that the memory stores executable code, and when the processor executes the executable code, the method described in the first aspect is implemented.
[0051] In the embodiments of this specification, a solution for training a user classification model is proposed. Using a dual classification model, clustering, classifying, and labeling pseudo-labels are respectively performed on the target domain data, and the pseudo-labels given by different classification models are comprehensively analyzed to select reliable users with unknown categories, and the classification models are trained based on these users. During the training process, both the sample clustering confidence and classification preference are considered, so that the classification model can, without accessing the source domain data, rely only on the target domain data to achieve the purpose of migrating to the target domain data, accurately classify the target domain data, and effectively identify the unknown categories existing in the target domain data. BRIEF DESCRIPTION OF THE DRAWINGS
[0052] In order to more clearly illustrate the technical solutions of the embodiments of the present invention, the drawings required for the description of the embodiments will be briefly introduced below. Obviously, the drawings in the following description are only some embodiments of the present invention, and those of ordinary skill in the art can also obtain other drawings based on these drawings without creative efforts.
[0053] Figure 1Schematic diagram of an implementation framework for training a user classification model provided according to an embodiment of this specification;
[0054] Figure 2 Schematic diagram of a method flow for training a user classification model disclosed in this specification;
[0055] Figure 3 Schematic diagram of an apparatus for training a user classification model provided according to an embodiment of this specification. Detailed implementation manners
[0056] In traditional machine learning theory, there is a core and important assumption: the independent and identically distributed (IID) assumption. This assumption includes two aspects: one is independence, which assumes that each sample in the dataset is independent of each other, that is, the occurrence of any one sample will not affect the generation of other samples. The other is identical distribution, which assumes that all samples are derived from the same probability distribution, meaning that each sample has the same statistical characteristics.
[0057] In short, the IID assumption states that the samples in the training set (source domain data) and the test set (target domain data) are independent of each other and follow the same probability distribution. This assumption provides an idealized framework for the training and application of machine learning models. In this idealized framework, it can be ensured that the training performance of the model on the source domain data can be effectively generalized to the target domain data that the model has not encountered. However, this idealized assumption often does not hold in the actual application scenarios of the model.
[0058] In practical applications, source domain data and target domain data usually exhibit different data characteristics due to factors such as environmental changes, time passage, and different data sources. For example, in the user classification scenario in the financial field, the source domain data may come from a certain financial business line, while the target domain data includes user data from all business lines in the enterprise, and there are differences in data characteristics and class probability distributions compared to the source domain data. This manifestation of difference can be called domain shift. Furthermore, with the expansion of the business, there may also be unknown class sample data in the target domain data that did not appear in the source domain data. This phenomenon can be called class shift.
[0059] After the model is trained based on the source domain data, it usually needs to be deployed to the target domain data to perform various downstream tasks. Compared with the source domain data, the target domain data often has the two types of offsets introduced above. To enable the model to maintain stable performance on the target domain data, these data offsets pose severe challenges to the generalization ability and openness of the model. Specifically, for the domain offset between the target domain data and the source domain data, the model needs to have sufficient generalization ability to ensure that it can make accurate predictions on a brand-new target domain data; for the class offset between the target domain data and the source domain data, the model needs to demonstrate sufficient openness to be able to identify sample data belonging to unknown classes on the target domain data, rather than misclassifying them as any known class or simply ignoring these newly emerging samples.
[0060] In practical applications, another challenge that is usually faced is the inaccessibility of the source domain data. As mentioned before, when migrating the model from the source domain data to the target domain data, due to factors such as data privacy protection and legal regulations, the source domain data cannot be accessed, which limits the use of labeled sample data by the model and brings great difficulties to the domain adaptation training of the model.
[0061] In view of this, the inventors propose a method for training a user classification model in one or more embodiments of the present invention. For a pre-trained classification model, this method can effectively train the classification model without accessing the source domain data, relying only on the unlabeled target domain data, enabling it to adapt to the target domain data, overcome domain offset and class offset, and not only accurately detect known class samples on the target domain data, but also discover unknown class samples.
[0062] Figure 1An implementation framework for training a user classification model is shown. In the embodiments of this specification, the source domain dataset used by the pre-trained model during the training phase is in an inaccessible state. For a user set (target domain data) containing several users (shown as hollow circles in the accompanying drawings), two classification models initialized by the pre-trained model are adopted to perform cluster prediction and classification prediction on the user set respectively. Among them, before performing cluster prediction, the first classification model first conducts cluster division training on the user set to learn the global cluster structure mining ability adapted to the user set. And the second classification model can perform classification prediction on the user set based on the knowledge of the pre-trained model. According to the prediction results of the two classification models respectively, the first users belonging to unknown categories are screened out from the user set. These first users can be used as training data to further train the first / second classification model. During this process, the first classification model is used to train the confidence in the prediction results of the first users, while the second classification model is used to train the ability to predict the first users as known categories. The two classification models complement each other and are co-trained. After multiple rounds of training like this, the first / second classification model can achieve a balance between generalization ability and generating moderate classification errors, realizing the goal of migrating the classification model to the target domain data, and effectively solving the application problem of the classification model in the target data domain when the source domain data is inaccessible.
[0063] The above is a general description of a method for training a user classification model provided by the embodiments of this specification. Next, the solutions provided by one or more embodiments in this specification will be described in detail with reference to the accompanying drawings.
[0064] Figure 2 A flowchart of a method for training a user classification model according to the embodiments of this specification is shown. It can be understood that this method can be executed by any device, equipment, platform, or device cluster with computing and processing capabilities. Such as Figure 2As shown, in this embodiment, the method at least includes the following steps: Step S201: Obtain a first classification model and a second classification model respectively initialized based on a pre-trained model, where the pre-trained model is trained based on a data set of M user groups. Step S203: Perform cluster partitioning on an unlabeled user set, and use the cluster partitioning result to perform a first update on the first classification model. Step S205: Use the first classification model after the first update to perform cluster prediction on each user in the first subset of the user set to obtain cluster pseudo-labels; and use the second classification model to perform classification prediction of the M user groups on each user to obtain classification pseudo-labels. Step S207: Determine a first user with an unknown category from the first subset according to the cluster pseudo-labels and classification pseudo-labels of each user. Step S209: Use the first / second classification model to respectively predict the augmented samples of the first user to obtain first / second prediction results; with the goal that the first prediction result is close to the cluster pseudo-label of the first user, update the first classification model; with the goal that the sum of the prediction probabilities corresponding to the M user groups in the second prediction result is maximized, update the second classification model.
[0065] The specific implementation manners of the above steps will be described in detail below.
[0066] First, in step S201, obtain a first classification model and a second classification model respectively initialized based on a pre-trained model, where the pre-trained model is trained based on a data set of M user groups.
[0067] Initializing the first / second classification model using a pre-trained model means that the training of the classification model does not start with randomly initialized parameters, but uses the parameters of a pre-trained and mature model as the starting point for training the classification model. Since the pre-trained model has been trained with a large amount of training data (source domain data), it can possess a certain ability to capture features. Therefore, the classification model initialized by the pre-trained model can possess a certain ability to mine domain features at the initial stage of training, which can significantly shorten the time required for model training. It should be understood that the domain feature mining ability possessed by the classification model at this time comes from the source domain data (i.e., the data set containing the M user groups).
[0068] Among the source domain data used by the pre-trained model during training, it contains user samples of M user groups, and the pre-trained model has the ability to accurately classify user samples into the corresponding user groups. The source domain data can be shown as where there are N s user samples, and the category space is In this embodiment,
[0069] The user set in the target domain can be represented as where there are N t users, and the class space is The class space of the unknown class is where, ·\· represents the difference set of two sets.
[0070] For the sake of clarity and simplicity, each network model can also be represented by a formula. The pre-trained model can be denoted as ψ(·), the first classification model can be denoted as φ o (·), and the second classification model can be denoted as φ c (·).
[0071] After initializing the first classification model and the second classification model, based on the target domain data, model migration preparation work can be carried out. In step S203, the unlabeled user set is clustered, and using the clustering result, the first classification model is updated for the first time.
[0072] As mentioned above, the source domain data set used by the pre-trained model in the training phase usually exhibits different data characteristics due to factors such as environmental changes, time lapse, and different data sources compared to the target domain user set, and there may be domain shift or class shift.
[0073] In a specific practice, the source domain data and the target domain data can be business data sets collected from different business fields respectively. That is to say, the data set comes from the first business domain, and the user set comes from the second business domain.
[0074] Therefore, if the classification model initialized by the pre-trained model is directly used to classify the target domain data, it may be affected by domain shift, resulting in biased classification results. Further, due to the possible unknown class users in the target domain data, they will become noise in the classification, thus further affecting the accuracy of classification prediction.
[0075] To overcome the above problems, especially the influence brought by domain shift, the class cluster mining ability of the first classification model can be trained first. Specifically, the first classification model can be used to cluster and identify the users in the user set according to the representation similarity of each user, and based on the clustering result, the first classification model is updated so that the first classification model has a certain global cluster structure mining ability for the user set. This is because, even if the target domain data has a domain shift compared to the source domain data, the users of the same class usually have higher representation similarity (i.e., are closer to each other) in the target feature space. This relative representation similarity between users can be discovered only by representation mining in the target domain data without relying on the source domain data.
[0076] According to one implementation method, in order to enable the first classification model to accurately extract characteristic information of users in the target domain data for accurate cluster division, data enhancement may be performed on the user set to obtain an expanded sample set before the first update; based on the user set and the expanded sample set, a comparative learning method may be used to train the first classification model.
[0077] Specifically, the contrastive learning method can train the first classification model to more accurately extract the user's representation information. In contrastive learning, the positive samples of any user are usually obtained by data enhancement, and the negative samples come from other users and their corresponding data enhanced samples. During the training process of the first classification model, it is necessary to shorten the representation distance of the positive sample pairs (user samples and their data enhanced samples) and push the representation distance of the negative sample pairs (user samples and other user samples or other user samples' data enhanced samples) further. This training mechanism can effectively train the first classification model to capture the representation differences between user samples and enhance its ability to extract and distinguish user representation information.
[0078] The following steps can be followed for contrastive learning training of the first classification model: First, for each user sample in the user set, a data enhancement algorithm (for example, RandAugment algorithm) is used to generate corresponding positive samples. At the same time, other user samples are selected from the user set, or data enhancement samples of other user samples are used as negative samples. Then, the first classification model is used to calculate the representation distance of each positive sample pair and negative sample pair in the representation space (which can be measured by, for example, Euclidean distance, representation vector similarity, etc.), and the contrast loss is determined based on the representation distance (which can be the average value of the representation distance). Finally, according to the calculated contrast loss, the gradient of the first classification model parameters is calculated by the back propagation algorithm, and the first classification model parameters are updated accordingly to minimize the contrast loss. The above training process is iteratively executed until the preset training rounds are met or the first classification model achieves satisfactory performance. Through contrastive learning training, the first classification model can more accurately extract the user representation information in the user set, thereby improving the accuracy of clustering the user set.
[0079] Next, we return to step S203, in which the unlabeled user set is firstly clustered, and the clustering can be performed using a variety of clustering algorithms, such as K-Means, DBSCAN, etc.
[0080] According to one implementation, a K-nearest neighbor graph can be constructed for the user set, which includes nodes representing users and connection edge weights representing the similarity between the nodes. Processing the K-nearest neighbor graph based on a community discovery algorithm yields several communities. The first K value is determined based on the evaluation of the first type of clusters of the communities obtained under different K values. The several communities included in the K-nearest neighbor graph constructed using the first K value are used as the cluster division result.
[0081] Specifically, first, for each user in the user set, based on the feature vectors extracted by the first classification model, its K nearest neighbor users are retrieved, and a K-nearest neighbor graph is constructed accordingly. It includes nodes representing users and connection edge weights representing the similarity between the nodes. It can be understood that in the K-nearest neighbor graph, the higher the similarity between samples, the closer the distance between the corresponding nodes, and vice versa. The K-nearest neighbor graph can be shown as:
[0082]
[0083] Among them, is the set of K nearest neighbors of user i, sim(·,·) is the cosine similarity function, and v i is the feature vector of user i.
[0084] Then, processing the K-nearest neighbor graph based on a community discovery algorithm yields several communities. The community discovery algorithm is a technique in graph theory used to partition the nodes in a graph into different communities (i.e., clusters). Nodes in the same community are more closely connected in the graph, while nodes between different communities are less connected. Since the K-nearest neighbor graph is constructed based on the similarity between users, nodes in the same community have a high similarity, while nodes between different communities have a low similarity. That is to say, there is a high degree of discrimination between user nodes in different communities.
[0085] Next, based on the evaluation of the first type of clusters in the communities obtained under different K values, the first K value is determined. Specifically, for each community obtained based on the community discovery algorithm, evaluate the quality of it as a cluster, which can be evaluated by various metrics. For example, the modularity of the cluster, the silhouette coefficient of the cluster, etc. Adjust the K value according to the evaluation results of the cluster quality. For example, if it is found that the connections between nodes within a certain community are not tight enough, or the distinguishability between different communities is not high, then it may be necessary to increase the K value to introduce more neighboring users and improve the community structure. On the contrary, if it is found that the connections between nodes within a certain community are too tight, or the distinguishability between different communities is too high, then it may be necessary to decrease the K value to avoid overfitting. This process of adjusting the K value usually needs to be iteratively executed for multiple rounds. In each round of iteration, according to the K value adjusted in the previous time, reconstruct the K-nearest neighbor graph, discover the communities therein, and perform a new round of evaluation and K value adjustment until the optimal K value (i.e., the first K value) is found.
[0086] In a specific example, an elbow graph can be constructed based on the K-nearest neighbor graph to determine an optimal K value (i.e., the first K value) in the elbow graph using the elbow method. In this example, the sum of squared errors SSE (i.e., in the K-nearest neighbor graph, the sum of the squares of the distances from each node to its nearest clustering center) can be used as the dependent variable (i.e., the vertical coordinate system of the elbow graph), and the corresponding K value as the independent variable (i.e., the horizontal coordinate system of the elbow graph) to construct the elbow graph. As the K value decreases, the user set will be divided more and more finely, the tightness of the nodes in each community will be higher and higher, and the SSE will also decrease accordingly. When the K value is greater than the optimal K value, gradually decreasing the K value will significantly reduce the SSE. When the K value reaches the optimal K value, further decreasing the K value, the rate of decrease of the SSE will suddenly slow down, showing an inflection point in the elbow graph. Therefore, the K value corresponding to the point presented as the "elbow" in the elbow graph can be determined as the first K value, and the several communities included in the K-nearest neighbor graph constructed with this K value are used as the cluster division result.
[0087] Through the above steps, not only the optimal cluster division result is determined, but also the rationality and distinguishability of each cluster in the feature space are ensured. It should be understood that in practical applications, different community discovery algorithms can be used to implement the above steps. For example, algorithms based on modularity optimization, algorithms based on spectral clustering, etc. Similarly, the first type of cluster evaluation can also be calculated by different evaluation methods.
[0088] In a specific example, the community discovery algorithm is the Louvain algorithm, and the first type of cluster evaluation is the within-cluster sum of squares evaluation (Cluster Sum of Square).
[0089] After determining the first K value (with k *After showing), the number of clusters in the corresponding K-nearest neighbor graph can be obtained, denoted as The cluster division result indicates the target cluster to which each user belongs. Based on the cluster division result, the first classification model can be first updated.
[0090] According to one implementation, the first update includes: performing cluster prediction on each user in the second subset of the user set to obtain the prediction vector of each user, including the prediction probability that the user belongs to each cluster. Comparing the prediction probability corresponding to the target cluster in the prediction vector with a preset first threshold, and according to the comparison result, dividing the users in the second subset into member users and non-member users. Taking the maximization of the sum of the first probability and the second probability as the goal, updating the first classification model; the first probability is positively correlated with the prediction probability of the member user on the corresponding target cluster; the second probability is positively correlated with the sum of the prediction vector similarities between the non-member user and its respective neighbor samples.
[0091] The first update can be iteratively executed for multiple rounds. Through multiple rounds of training, the first classification model can exhibit better global cluster structure mining ability on the user set. In addition, the training of each round can be performed on the second subset of the user set according to a predetermined hyperparameter (which can be the number of training samples).
[0092] The first threshold can be a preset hyperparameter for evaluating the relative relationship between the user and the clustering prototype (for example, it can be the lowest prediction probability). The clustering prototype can be the centroid of each cluster. Each clustering prototype serves as the representative of its corresponding cluster. Users close to the clustering prototype can be considered as member users of the cluster, and vice versa can be considered as non-member users. The first threshold is denoted as γ. Member users can be represented as Non-member users can be represented as where p i,o is the softmax output of user i on the first classification model, is p i,o corresponding to the prediction probability of the user on cluster m. Thus, the user set composed of member users can be obtained and the user set composed of non-member users
[0093] Next, for the first classification model, the following loss function is used for training:
[0094]
[0095] where I i is the nearest neighbor index set of user (i.e., each neighbor sample), is the one-hot version (one-hot encoding), and are respectively and the m-th element of. B represents all users in this training round (i.e., the second subset), is a regularization term that helps the first classification model balance between improving generalization ability and generating appropriate training errors, preventing model collapse.
[0096] It can be seen that in the first term of this loss function, for each user (member user) in the user set , represents taking its predicted probability on the corresponding target cluster ( only has a value of 1 on the corresponding target cluster, and the values of the remaining elements are 0), so it is not difficult to understand that the first term represents the predicted probabilities of each member user on the corresponding target cluster. In the second term of this loss function, for each user (non-member user) in the user set , represents taking the sum of the similarity between the predicted vectors of this non-member user and each of its neighbor samples r. Generally speaking, it can be known that in this training, the first classification model aims to maximize the sum of the first probability and the second probability to update the model parameters (in specific operations, the parameters can be updated and adjusted by means of gradient descent, etc.). The first probability is positively correlated with the predicted probability of the member user on the corresponding target cluster; the second probability is positively correlated with the sum of the similarity between the predicted vectors of the non-member user and each of its neighbor samples.
[0097] After the first update of the first classification model, next, in step S205, using the first classification model after the first update, perform cluster prediction on each user in the first subset of the user set to obtain cluster pseudo-labels; and, using the second classification model, perform classification prediction on each user for the M user groups to obtain classification pseudo-labels.
[0098] The updated first classification model can be used to perform cluster prediction on each user in the first subset of the user set and assign corresponding cluster pseudo-labels to each user
[0099] The first subset can be obtained on the user set according to pre-determined hyperparameters (for example, the number of training samples).
[0100] As described above, the source domain dataset used by the pre-trained model in the training phase may have class bias compared to the target domain user set. Therefore, using the second classification model initialized by the pre-trained model to classify and predict the target domain data may be affected by class bias, resulting in deviation in class division, which can be manifested as the deviation of the classification boundary or the inability to effectively identify new unknown classes in the target domain data.
[0101] To overcome the above problems, a statistical method can be adopted. According to the class probability distribution characteristics existing in the target domain data, a threshold interval for classification decision is statistically obtained on the prediction output of the second classification model. This is because, although the class attribution of a single user may be affected by class bias, from the perspective of the classification prediction statistical characteristics of the overall users, the relative probability distribution between different classes usually tends to be stable. By determining the threshold intervals corresponding to these classes respectively, the classification boundary can be more accurately defined, thereby improving the classification performance of the second classification model on the target domain data. Even when facing new unknown classes, it can also maintain a high recognition ability.
[0102] According to one implementation, the classification prediction results of each user output by the second classification model can be obtained, and each user is classified into the group set of the user group corresponding to the highest prediction probability in the classification prediction results. Based on the group set, the attribution threshold and exclusion threshold corresponding to the user group are determined. For any user in any group set, if its highest prediction probability is higher than the attribution threshold, the classification pseudo-label of this user is marked as the user group corresponding to this group set; if its highest prediction probability is lower than the exclusion threshold, the classification pseudo-label of this user is marked as an unknown class.
[0103] In this implementation, a second classification model can be used to predict each user in the user set to obtain the prediction probability of each user corresponding to each category. Each user is classified into the user group corresponding to the highest prediction probability in its prediction result. In this way, each user group constitutes its own group set. According to the user prediction probabilities in each group set, a classification threshold (such as a probability threshold, percentile, etc.) and an exclusion threshold (such as a probability threshold, percentile, etc.) can be set for each user group. For example, the median of the prediction probabilities of all users in the user group on this user group is statistically calculated as the classification threshold. Similarly, the exclusion threshold is determined. Generally, the exclusion threshold and the classification threshold can form a continuous probability interval, but in some practices, there may also be a certain intersection or not be completely complementary between the interval formed by the exclusion threshold and the interval formed by the classification threshold, and no specific limitation is made here. After determining the classification threshold and the exclusion threshold, each user in the group set can be classified according to these two thresholds and the corresponding classification pseudo-label is marked. Specifically, if the highest prediction probability of a user is higher than the attribution threshold, the classification pseudo-label of this user is marked as the user group corresponding to this group set; if the highest prediction probability of a user is lower than the exclusion threshold, the classification pseudo-label of this user is marked as an unknown category.
[0104] In a specific example, determining the attribution threshold and the exclusion threshold for the corresponding user group may include: for any group set, fitting the attribution probability distribution and the exclusion probability distribution based on the maximum likelihood estimation algorithm. According to the expected probabilities of the attribution / exclusion probability distributions, the attribution threshold and the exclusion threshold for the user group corresponding to this group set are respectively determined.
[0105] Specifically, first use the second classification model to perform classification prediction on each user in the first subset. In the softmax processing of the output layer, a preset temperature coefficient can be used for calibration, which is not limited here. The softmax output corresponding to the user is denoted as q i . For any user group d in the M user groups, the user with the highest prediction probability on this user group is included in the group set Q of this user group d . In this way, the group set corresponding to each user group d can be obtained, which contains the users with the highest prediction probability confidence on d. For any group set Q d , it can be considered that the probabilities of the users in it belonging to this user group all conform to an independent probability distribution. Therefore, the maximum likelihood estimation algorithm can be used to fit the attribution probability distribution of belonging to this user group and the exclusion probability distribution of not belonging to this user group on each group set, and then the corresponding attribution threshold and exclusion threshold are determined according to the expected probabilities of each probability distribution.
[0106] In a specific example, the maximum likelihood estimation algorithm is the Expectation-Maximization (EM) algorithm, and the membership probability distribution and the exclusion probability distribution are beta distributions. For any population set Q d , fit a binary beta mixture model through the EM algorithm and obtain the parameters of the fitted mixture model The membership threshold corresponding to the user population d Exclusion threshold can be taken from the mean value (mathematical expectation probability) of the binary beta mixture model, that is Next, according to the membership threshold and the exclusion threshold corresponding to each user population, for each user in the first subset perform classification and mark the corresponding classification pseudo-label It can be expressed as:
[0107]
[0108] Through the above steps, using the first classification model and the second classification model respectively, each user in the first subset of the user set is labeled with a cluster pseudo-label Classification pseudo-label Next, in step S207, according to the cluster pseudo-labels and classification pseudo-labels of each user, determine the first user with an unknown category from the first subset.
[0109] In this step, the pseudo-labels output by the two classification models are consensus / associated to obtain a more accurate user classification. According to one implementation, the classification pseudo-label may include: an unknown category label and a known category label representing one of the M user populations. In this step, determining the first user with an unknown category from the first subset may include: determining the target mapping relationship between the cluster pseudo-label and the M user populations, so that the number of consensus users whose cluster pseudo-label and classification pseudo-label conform to the target mapping relationship is maximized. Determine the set of users with unknown categories, including users with a classification pseudo-label of unknown category label and users with a classification pseudo-label of known category label but not belonging to the consensus users; the first user is selected from the set of users with unknown categories.
[0110] The set composed of users with a cluster pseudo-label of u can be shown as: The set composed of users belonging to the user population d can be shown as: The set composed of consensus users identified by the target mapping relationship can be defined as: The process of determining the target mapping relationship is to find the target mapping relationship that can maximize the total number of users in each set composed of consensus users.
[0111] According to one implementation, a matching matrix can be constructed based on class cluster pseudo-labels and classification pseudo-labels. On the matching matrix, the Hungarian algorithm is used to find the matching relationship that can maximize the total number of consensus users from a global perspective, which is the target mapping relationship. Specifically, the matching matrix wherein, is the number of known user groups; is the number of clusters represented by the class cluster pseudo-labels. Any matrix element a in the matching matrix du takes the value of the negative of the number of consensus users. The purpose is to cooperate with the execution of the Hungarian algorithm. The Hungarian algorithm aims to find the minimum total cost matching in the matching matrix. Therefore, the value of the matrix element is taken from the negative of the number of consensus users, which means that the optimization problem of finding the target mapping relationship is transformed into a minimization problem, and the Hungarian algorithm is just suitable for solving this type of optimization problem. After the Hungarian algorithm finishes execution, the obtained minimum total cost matching result is the target mapping relationship that maximizes the sum of the number of consensus users.
[0112] In this way, under the target mapping relationship, each mapping set The set composed of users belonging to user group d: where The set composed of all consensus users can be denoted as In the definition of the set of users with unknown categories, it can be considered that during the consensus process, each user not associated by the target mapping relationship belongs to the users who are confirmed to belong to a certain user group on the second classification model but cannot find a corresponding suitable cluster on the first classification model. These users who cannot reach a consensus can be regarded as users with unknown categories. Therefore, the set of users with unknown categories can be defined as users with the classification pseudo-label of unknown category label (unknown) and users with the classification pseudo-label of known category label but not belonging to the consensus users, that is:
[0113] Based on the set of users with unknown categories obtained after the above model consensus, the first users included therein can be regarded as reliable users with unknown categories. Using these first users as training data, the first classification model and the second classification model are trained respectively, which can improve the accuracy of the model in detecting unknown categories on the target domain data. Therefore, in step S209, the first / second classification model is used to predict the augmented samples of the first user respectively to obtain the first / second prediction results; the first classification model is updated with the goal that the first prediction result is close to the class cluster pseudo-label of the first user; the second classification model is updated with the goal that the sum of the prediction probabilities corresponding to the M user groups in the second prediction result is maximized.
[0114] In this step, training can be performed on the first classification model with the goal of enhancing its prediction confidence for the first user. Therefore, the first classification model can be used to predict the enhanced samples of the first user, and the first classification model can be updated with the goal that the prediction result is close to the cluster pseudo-label of the first user. The loss function used can be defined as:
[0115]
[0116] where represents the m-th element of the softmax output of the corresponding enhanced sample on the first classification model after the user represents the cluster pseudo-label of the user and is the m-th element of the one-hot version (one-hot encoding) of
[0117] Meanwhile, training can be performed on the second classification model with the goal of making its prediction result for the first user approach any user group. Therefore, the second classification model can be used to predict the enhanced samples of the first user, and the second classification model can be updated with the goal of maximizing the sum of the prediction probabilities corresponding to the M user groups in the prediction result. The loss function used can be defined as:
[0118]
[0119] where represents the m-th element of the softmax output of the corresponding enhanced sample on the second classification model after the user
[0120] The above is an introduction to the main process of a method for training a user classification model provided in the embodiments of this specification. Although the above embodiments mainly use the user classification task as an example to elaborate on the method process. However, the technical concept embodied therein can also be applied to the training process of classification models for other similar tasks.
[0121] In addition, in some embodiments of this specification, a method for training a classification model using consensus users is also provided. According to one implementation, the first / second classification model is used to predict the enhanced samples of any consensus user respectively to obtain the first / second consensus item prediction results; the first / second classification model is updated respectively with the goal that the first / second consensus item prediction results are close to the cluster / classification pseudo-label of the consensus user.
[0122] Consensus users belong to two classification models. At the level of cluster / classification detection, users for which consensus can be reached. During training, this consensus can be strengthened so that the classification model has a higher confidence in detecting the categories of consensus users. Therefore, enhanced samples of consensus users can be used as positive samples, and the classification model is used for prediction. The model is updated with the goal that the prediction results of the enhanced samples are consistent with the prediction results of the consensus users.
[0123] Corresponding to the above training objective, the loss function can be defined as:
[0124] The first classification model:
[0125]
[0126] Among them, represents the m-th element of the softmax output of the corresponding enhanced sample on the first classification model after the user represents the cluster pseudo-label of the user in the one-hot version (one-hot encoding) of the
[0127] The second classification model:
[0128]
[0129] Among them, represents the m-th element of the softmax output of the corresponding enhanced sample on the second classification model after the user represents the classification pseudo-label of the user in the one-hot version (one-hot encoding) of the
[0130] In some other embodiments of this specification, a method for training a classification model using a user to be determined is also provided. The sample to be determined includes other users in the first subset that do not belong to the consensus users and do not belong to the set of users with unknown categories. The set of users to be determined is marked with . During training, for any user to be determined, the prediction result vectors of the user to be determined and its respective neighboring samples output by the first classification model are obtained, and the first classification model is updated with the goal of maximizing the similarity between the prediction result vectors. The corresponding loss function can be defined as:
[0131]
[0132] Among them, p i,o represents the Softmax output on the first classification model. I i is the set of nearest neighbor indices of the user (i.e., each neighboring sample). B represents all users in the current training round (the first subset), is a regularization term to prevent model collapse.
[0133] Meanwhile, for the second classification model, the second classification model can be used to predict the strongly augmented and weakly augmented samples of any user to be determined respectively; and with the goal of maximizing the consistency of the prediction results, the second classification model is updated. The corresponding loss function can be defined as:
[0134]
[0135] where and represent the m-th element of the softmax output on the second classification model of the strongly augmented sample and the weakly augmented sample corresponding to the user after strong augmentation and weak augmentation respectively.
[0136] The above describes in detail a method for training a user classification model according to one or more embodiments. By using the above method provided in the embodiments of the present specification, without the need to access the source domain data of the pre-trained model, only relying on the unlabeled target domain data, the pre-trained model can be effectively trained into a classification model adapted to the target domain data; enabling it to overcome domain shift and class shift, and not only being able to accurately detect known class samples on the target domain data, but also being able to discover unknown class samples.
[0137] In this specification, the "first" in terms such as the first classification model and the first probability distribution, and the corresponding "second" (if any) in the text are only for the convenience of distinction and description, and do not have any restrictive meaning.
[0138] The above content describes specific embodiments of this specification, and other embodiments are within the scope of the appended claims. In some cases, the actions or steps recited in the claims can be performed in a different order from that in the embodiments, and still achieve the desired results. Additionally, the processes depicted in the drawings do not necessarily have to be performed in the specific order or continuous order shown to achieve the desired results. In certain embodiments, multitasking and parallel processing are also possible, or may be advantageous.
[0139] Figure 3Schematic diagram of an apparatus for training a user classification model according to an embodiment of this specification. The apparatus 300 is deployed in a computing device, which can be implemented by any device, equipment, platform, device cluster, etc. with computing and processing capabilities. This apparatus embodiment corresponds to Figure 2 the method embodiment shown. The apparatus 300 includes:
[0140] An acquisition module 301, configured to acquire a first classification model and a second classification model respectively initialized based on a pre-trained model, where the pre-trained model is trained based on a data set of M user groups.
[0141] A first update module 302, configured to perform cluster partitioning on an unlabeled user set, and use the cluster partitioning result to perform a first update on the first classification model.
[0142] A prediction module 303, configured to use the first classification model after the first update to perform cluster prediction on each user in the first subset of the user set to obtain cluster pseudo-labels; and use the second classification model to perform classification prediction on each user for the M user groups to obtain classification pseudo-labels.
[0143] A determination module 304, configured to determine a first user of an unknown category from the first subset according to the cluster pseudo-labels and classification pseudo-labels of each user.
[0144] A second update module 305, configured to use the first / second classification model to respectively predict the augmented samples of the first user to obtain first / second prediction results; aiming at the first prediction result being close to the cluster pseudo-label of the first user, update the first classification model; aiming at maximizing the sum of the prediction probabilities corresponding to the M user groups in the second prediction result, update the second classification model.
[0145] According to an embodiment of another aspect, this specification also provides a computer program product, including a computer program / instructions, which when executed by a processor, implement the steps of the foregoing method in combination with Figure 2 the method described.
[0146] According to an embodiment of yet another aspect, this specification also provides a computing device, including a memory and a processor, characterized in that an executable code is stored in the memory, and when the processor executes the executable code, the steps of the foregoing method in combination with Figure 2 the method described are implemented.
[0147] Those skilled in the art should be able to realize that in one or more of the above examples, the functions described in the embodiments of the present invention can be implemented by hardware, software, firmware, or any combination thereof. When implemented using software, these functions can be stored in a computer-readable medium or transmitted as one or more instructions or codes on a computer-readable medium.
[0148] The specific embodiments described above have further elaborated on the objectives, technical solutions, and beneficial effects of the embodiments of the present invention. It should be understood that the above is only the specific embodiments of the embodiments of the present invention and is not used to limit the protection scope of the present invention. Any modifications, equivalent replacements, improvements, etc. made on the basis of the technical solutions of the present invention shall be included in the protection scope of the present invention.
Claims
1. A method for training a user classification model, comprising: Obtaining a first classification model and a second classification model respectively initialized based on a pre-trained model, wherein the pre-trained model is trained based on a data set of M user groups; Clustering the unlabeled user set, and using the clustering result to perform a first update on the first classification model; Using the first updated first classification model, performing cluster prediction on each user in the first subset of the user set to obtain a cluster pseudo label; And, using the second classification model, performing classification prediction of the M user groups on each of the users to obtain classification pseudo labels; Determine a first user of unknown category from the first subset according to the cluster pseudo label and the classification pseudo label of each user; Using the first / second classification model, respectively predict the enhanced sample of the first user to obtain a first / second prediction result; The first classification model is updated with the goal of making the first prediction result close to the cluster pseudo label of the first user; the second classification model is updated with the goal of maximizing the sum of the prediction probabilities corresponding to the M user groups in the second prediction result.
2. The method according to claim 1, wherein: The data set comes from a first business domain, and the user set comes from a second business domain.
3. The method according to claim 1, wherein: Before the first updating, the method further includes: Performing data enhancement on the user set to obtain an expanded sample set; and training the first classification model by using a contrastive learning method based on the user set and the expanded sample set.
4. The method according to claim 1, wherein: The clustering of the unlabeled user set includes: For the user set, a K-nearest neighbor graph is constructed, which includes nodes representing users and connection edge weights representing similarities between representative nodes; Processing the K nearest neighbor graph based on a community discovery algorithm to obtain several communities; Determine the first K value based on the first cluster evaluation of the community obtained under different K values; Several communities included in the K nearest neighbor graph constructed by using the first K value are used as the clustering result.
5. The method according to claim 4, wherein: The community discovery algorithm is the Louvain algorithm, and the first type of cluster evaluation is an intra-cluster square sum evaluation.
6. The method according to claim 1, wherein: The clustering result indicates the target cluster to which each user belongs; the first updating includes: Performing cluster prediction on each user in the second subset of the user set to obtain a prediction vector for each user, including a prediction probability that the user belongs to each cluster; Comparing the predicted probability corresponding to the target cluster in the prediction vector with a preset first threshold, and dividing the users in the second subset into member users and non-member users according to the comparison result; The first classification model is updated with the goal of maximizing the sum of the first probability and the second probability; the first probability is positively correlated to the predicted probability of the member user in the corresponding target cluster; the second probability is positively correlated to the sum of the similarities of the predicted vectors between the non-member user and each of its neighboring samples.
7. The method according to claim 1, wherein: The obtaining of the classification pseudo label comprises: Obtaining the classification prediction result of each user output by the second classification model, and classifying each user into a group set of user groups corresponding to the highest prediction probability in the classification prediction result; Based on the group set, determine the attribution threshold and exclusion threshold of the corresponding user group; For any user in any group set, if its highest predicted probability is higher than the attribution threshold, the classification pseudo label of the user is marked as the user group corresponding to the group set; if its highest predicted probability is lower than the exclusion threshold, the classification pseudo label of the user is marked as an unknown category.
8. The method according to claim 7, wherein: The determining of the attribution threshold and the exclusion threshold of the corresponding user group based on the group set includes: For any group set, based on the maximum likelihood estimation algorithm, fit the probability distribution of belonging and the probability distribution of exclusion; According to the expected probability of the attribution / exclusion probability distribution, the attribution threshold and the exclusion threshold of the user group corresponding to the group set are determined respectively.
9. The method according to claim 8, wherein: The maximum likelihood estimation algorithm is the Expected Maximum (EM) algorithm, and the attribution probability distribution and the exclusion probability distribution are Beta distributions.
10. The method according to claim 1, wherein: The classification pseudo-labels include unknown category labels and known category labels representing one of the M user groups; Determining a first user of an unknown category from the first subset includes: Determine the target mapping relationship between cluster pseudo labels and M user groups, so as to maximize the number of consensus users whose cluster pseudo labels and classification pseudo labels conform to the target mapping relationship; An unknown category user set is determined, which includes users whose classification pseudo labels are unknown category labels and users whose classification pseudo labels are known category labels but are not consensus users; the first user is selected from the unknown category user set.
11. The method according to claim 10, wherein: The method further comprises: The first / second classification model is used to predict the enhanced samples of any consensus user to obtain the prediction results of the first / second consensus item; the first / second classification model is updated with the goal that the prediction results of the first / second consensus item are close to the cluster / classification pseudo-label of the consensus user.
12. The method according to claim 10, wherein: The method further comprises: For any user to be determined, obtain the prediction result vector of the user to be determined and each of its neighboring samples output by the first classification model, and update the first classification model with the goal of maximizing the similarity between the prediction result vectors; the user to be determined includes other users in the first subset who are neither consensus users nor unknown category users; The second classification model is used to predict the samples of any to-be-determined user after strong enhancement and weak enhancement, respectively; and the second classification model is updated with the goal of maximizing the consistency of the prediction results.
13. A device for training a user classification model, comprising: An acquisition module is configured to acquire a first classification model and a second classification model respectively initialized based on a pre-trained model, wherein the pre-trained model is trained based on a data set of M user groups; A first updating module is configured to perform clustering on the unlabeled user set and perform a first update on the first classification model using the clustering result; A prediction module is configured to use the first updated first classification model to perform cluster prediction on each user in the first subset of the user set to obtain a cluster pseudo label; And, using the second classification model, performing classification prediction of the M user groups on each of the users to obtain classification pseudo labels; A determination module configured to determine a first user of an unknown category from the first subset according to the cluster pseudo-label and the classification pseudo-label of each user; A second updating module is configured to use the first / second classification model to predict the enhanced sample of the first user respectively to obtain a first / second prediction result; The first classification model is updated with the goal of making the first prediction result close to the cluster pseudo label of the first user; the second classification model is updated with the goal of maximizing the sum of the prediction probabilities corresponding to the M user groups in the second prediction result.
14. A computer program product, comprising a computer program / instruction, which, when executed by a processor, implements the steps of the method according to any one of claims 1 to 12.
15. A computing device comprising a memory and a processor, characterized in that: The memory stores executable codes, and when the processor executes the executable codes, the method according to any one of claims 1 to 12 is implemented.
Citation Information
Cited By
Defect classification model establishing method and defect classification method
CN122336446A