Cross-customer weighted federated domain adaptation for event detection in warehouses
One example method includes registering, by a customer, with a service provider, receiving, by the customer from the service provider, a global machine learning model, running, by the customer, the global machine learning model as a local machine learning model, collecting, by the customer, unlabeled data generated by edge devices operating in a customer domain, checking, by the customer, to determine if the customer domain has changed, and when it is determined that the customer domain has changed, performing, by the customer, a model adaptation process on the local machine learning model, and transmitting to the service provider, by the customer, gradients that comprise customer implemented changes to the local machine learning model.
1 . A method, comprising:
registering, by a customer computing system, with a service provider;
receiving, by the customer computing system from the service provider, a global machine learning model comprising a domain-adversarial neural network (DANN) including a feature extractor and a domain classifier;
running, by the customer computing system, the global machine learning model as a local machine learning model on edge devices operating in a customer domain;
collecting, by the customer computing system, unlabeled data generated by the edge devices operating in the customer domain;
extracting, by the feature extractor of the DANN executed by the customer computing system, features from the collected unlabeled data;
measuring a similarity between a probability distribution of extracted features of original data from an original data domain and a probability distribution of the extracted features of the collected unlabeled data, by the domain classifier of the customer computing system, to determine if the collected unlabeled data is not included in the customer domain based on the similarity;
when it is determined that the collected unlabeled data is not included in the customer domain, performing, by the customer computing system, a model adaptation process by applying a gradient in the customer domain and a gradient in a new domain of the collected unlabeled data to the global machine learning model to generate and update a local machine learning model from the global machine learning model; and
transmitting to the service provider, by the customer computing system, information indicating an amount of the collected unlabeled data by the customer computing system and gradients that comprise changes between the global machine learning model and the updated local machine learning model and comprise the gradient in the customer domain and the gradient in the new domain,
wherein gradients from the edge devices are used to update the global machine learning model by the service provider.
2 . The method as recited in claim 1 , wherein the global machine learning model comprises an event detection model operable to detect occurrence of specified events in a customer warehouse.
3 . The method as recited in claim 1 , wherein the global machine learning model comprises respective weighted gradients generated by the customer computing system, and by other customer computing systems.
4 . The method as recited in claim 1 , wherein the data collected by the customer remains confidential with the customer computing system and is not shared with other customer computing systems.
5 . The method as recited in claim 1 , wherein updating the local machine learning model comprises aggregating gradients and using the aggregated gradients to update the local machine learning model.
6 . The method as recited in claim 1 , wherein the customer computing system receives, from the service provider, a set of training data, and the training data is used to determine if the customer domain has changed.
7 . The method as recited in claim 1 , wherein determining if the collected unlabeled data is not included in the customer domain comprises comparing training data with the collected unlabeled data and obtaining a divergence between the training data and the collected unlabeled data.
8 . The method as recited in claim 1 , wherein the customer domain is deemed as having changed when a divergence between training data and the collected unlabeled data equals or exceeds a specified threshold.
9 . The method as recited in claim 1 , wherein the global machine learning model is received by the customer computing system from the service provider as-a-Service.
10 . The method as recited in claim 1 , wherein the global machine learning model is an updated version of the local machine learning model running at the customer computing system prior to receipt, by the customer computing system, of the global machine learning model.
11 . A non-transitory storage medium having stored therein instructions that are executable by one or more hardware processors to perform operations comprising:
registering, by a customer computing system, with a service provider;
receiving, by the customer computing system from the service provider, a global machine learning model comprising a domain-adversarial neural network (DANN) including a feature extractor and a domain classifier;
running, by the customer computing system, the global machine learning model as a local machine learning model on edge devices operating in a customer domain;
collecting, by the customer computing system, unlabeled data generated by edge devices operating in the customer domain;
extracting, by the feature extractor of the DANN executed by the customer computing system, features from the collected unlabeled data;
measuring a similarity between a probability distribution of extracted features of original data from an original data domain and a probability distribution of the extracted features of the collected unlabeled data, by the domain classifier of the customer computing system, to determine if the collected unlabeled data is not included in the customer domain based on the similarity;
when it is determined that the collected unlabeled data is not included in the customer domain, performing, by the customer computing system, a model adaptation process by applying a gradient in the customer domain and a gradient in a new domain of the collected unlabeled data to the global machine learning model to generate and update a local machine learning model from the global machine learning model; and
transmitting to the service provider, by the customer computing system, information indicating an amount of the collected unlabeled data by the customer computing system and gradients that comprise changes between the global machine learning model and the updated local machine learning model and comprise the gradient in the customer domain and the gradient in the new domain,
wherein gradients from the edge devices are used to update the global machine learning model by the service provider.
12 . The non-transitory storage medium as recited in claim 11 , wherein the global machine learning model comprises an event detection model operable to detect occurrence of specified events in a customer warehouse.
13 . The non-transitory storage medium as recited in claim 11 , wherein the global machine learning model comprises respective weighted gradients generated by the customer computing system, and by other customer computing systems.
14 . The non-transitory storage medium as recited in claim 11 , wherein the data collected by the customer remains confidential with the customer computing system and is not shared with other customer computing systems.
15 . The non-transitory storage medium as recited in claim 11 , wherein updating the local machine learning model comprises aggregating gradients and using the aggregated gradients to update the local machine learning model.
16 . The non-transitory storage medium as recited in claim 11 , wherein the customer computing system receives, from the service provider, a set of training data, and the training data is used to determine if the customer domain has changed.
17 . The non-transitory storage medium as recited in claim 11 , wherein determining if the collected unlabeled data is not included in the customer domain comprises comparing training data with the collected unlabeled data and obtaining a divergence between the training data and the collected unlabeled data.
18 . The non-transitory storage medium as recited in claim 11 , wherein the customer domain is deemed as having changed when a divergence between training data and the collected unlabeled data equals or exceeds a specified threshold.
19 . The non-transitory storage medium as recited in claim 11 , wherein the global machine learning model is received by the customer computing system from the service provider as-a-Service.
20 . The non-transitory storage medium as recited in claim 11 , wherein the global machine learning model is an updated version of a local machine learning model running at the customer computing system prior to receipt, by the customer computing system, of the global machine learning model.