Systems and methods for node weighting and aggregation for federated learning techniques
A system described herein may provide a technique for enhanced federated learning in an environment that makes use of one or more centralized models. Different nodes may be associated with different groups. Each node may provide refinement information for a given centralized model. The modifications for particular groups may be aggregated and the model may be modified based on modifications associated with each group, as opposed to modifications associated with each node. Weights for each group may be determined based on attributes of the modifications associated with each group, which may allow for the identification, on a group basis, of bias, maliciously injected data, outliers, and/or other types of modifications which may reduce the quality of the model. As such, embodiments described herein may enhance the quality, accuracy, and predictive ability of federated learning techniques that utilize distributed or federated modifications to a centralized model.
1 . A device, comprising:
one or more processors configured to:
receive policy information specifying respective sets of criteria for a plurality of different node groups, wherein the respective sets of criteria include a device type criteria;
identify that a first set of nodes is associated with a first device type;
identify, based on a first device type criteria specified by the policy information and further based on identifying that the first set of nodes is associated with the first device type, that the first set of nodes are associated with a first node group of the plurality of node groups;
identify that a second set of nodes is associated with a second device type;
identify, based on a second device type of criteria specified by the policy information and further based on identifying that the second set of nodes is associated with the second device type, that the second set of nodes are associated with a second node group of the plurality of node groups;
receive, from the first set of nodes, a first set of refinement information for a particular machine learning model, wherein the first set of refinement information includes a first plurality of modifications to a particular feature of the particular machine learning model, wherein each modification to the particular feature, of the first set of refinement information, has been generated by a respective node of the first set of nodes based on sensor data detected by the respective node of the first set of nodes;
receive, from the second set of nodes, a second set of refinement information for the particular machine learning model, wherein the second set of refinement information includes a second plurality of modifications to the particular feature of the particular machine learning model, wherein each modification to the particular feature, of the second set of refinement information, has been generated by a respective node of the second set based on sensor data detected by the respective node of the second set of nodes;
generate first aggregated refinement information based on first set of refinement information for the particular machine learning model;
generate second aggregated refinement information based on the second set of refinement information for the particular machine learning model;
determine a first weight associated with the first node group, wherein the first weight is determined based on identifying that the first node group is associated with the first device type;
determine a second weight associated with the second node group, wherein the second weight is determined based on identifying that the second node group is associated with the second device type;
identify a measure of error associated with the second set of refinement information, wherein identifying the measure of error includes:
applying the second set of refinement information to a set of data that includes a plurality of labels, and
identifying the measure of error based on applying the second set of refinement information to a particular subset, of the set of data, that is associated with a particular label of the plurality of labels;
identify that the measure of error associated with the second set of refinement information exceeds a threshold measure of error;
determine a third weight associated with the second group based on identifying that the measure of error, associated with the second refinement information, exceeds the threshold measure of error; and
modify the particular machine learning model based on the first and second aggregated refinement information and further based on the first and third weights.
2 . The device of claim 1 , wherein determining the first weight associated with the first node group includes:
determining a measure of variation between the refinement information associated with the first node group and the second node group.
3 . The device of claim 2 , wherein a first measure of variation between the first node group and the second node group is associated with the first weight for the first node group, and
wherein a second measure of variation, that is lower than the first measure of variation, between the first node group and the second node group is associated with a third weight, that is higher than the first weight, for the first node group.
4 . The device of claim 1 , wherein modifying the particular machine learning model based on the aggregated refinement information and the first and second weights respectively associated with the first and second node groups consumes fewer processing resources than modifying the particular machine learning model based on the plurality of instances of refinement information.
5 . The device of claim 1 , wherein the refinement information includes feature importance gradients for one or more features associated with the particular machine learning model.
6 . The device of claim 1 , wherein the measure of error associated with the second set of refinement information is a second measure of error, wherein determining the first weight associated with the first node group includes:
determining a first measure of error with respect to a particular label associated with a particular instance of refinement information associated with the first node group, wherein the first weight associated with the first node group is further based on the determined first measure of error.
7 . The device of claim 1 , wherein the first set of nodes includes a plurality of Internet of Things (“IoT”) devices that each include one or more respective sensors that detect the sensor data associated with each respective node of the first set of nodes.
8 . A non-transitory computer-readable medium, storing a plurality of processor-executable instructions to:
receive policy information specifying respective sets of criteria for a plurality of different node groups, wherein the respective sets of criteria include a device type criteria;
identify that a first set of nodes is associated with a first device type;
identify, based on a first device type criteria specified by the policy information and further based on identifying that the first set of nodes is associated with the first device type, that the first set of nodes are associated with a first node group of the plurality of node groups;
identify that a second set of nodes is associated with a second device type;
identify, based on a second device type of criteria specified by the policy information and further based on identifying that the second set of nodes is associated with the second device type, that the second set of nodes are associated with a second node group of the plurality of node groups;
receive, from the first set of nodes, a first set of refinement information for a particular machine learning model, wherein the first set of refinement information includes a first plurality of modifications to a particular feature of the particular machine learning model, wherein each modification to the particular feature, of the first set of refinement information, has been generated by a respective node of the first set of nodes based on sensor data detected by the respective node of the first set of nodes;
receive, from the second set of nodes, a second set of refinement information for the particular machine learning model, wherein the second set of refinement information includes a second plurality of modifications to the particular feature of the particular machine learning model, wherein each modification to the particular feature, of the second set of refinement information, has been generated by a respective node of the second set based on sensor data detected by the respective node of the second set of nodes;
generate first aggregated refinement information based on first set of refinement information for the particular machine learning model;
generate second aggregated refinement information based on the second set of refinement information for the particular machine learning model;
determine a first weight associated with the first node group, wherein the first weight is determined based on identifying that the first node group is associated with the first device type;
determine a second weight associated with the second node group, wherein the second weight is determined based on identifying that the second node group is associated with the second device type;
identify a measure of error associated with the second set of refinement information, wherein identifying the measure of error includes:
applying the second set of refinement information to a set of data that includes a plurality of labels, and
identifying the measure of error based on applying the second set of refinement information to a particular subset, of the set of data, that is associated with a particular label of the plurality of labels;
identify that the measure of error associated with the second set of refinement information exceeds a threshold measure of error;
determine a third weight associated with the second group based on identifying that the measure of error, associated with the second refinement information, exceeds the threshold measure of error; and
modify the particular machine learning model based on the first and second aggregated refinement information and further based on the first and third weights.
9 . The non-transitory computer-readable medium of claim 8 , wherein determining the first weight associated with the first node group includes:
determining a measure of variation between the refinement information associated with the first node group and the second node group.
10 . The non-transitory computer-readable medium of claim 9 , wherein a first measure of variation between the first node group and the second node group is associated with the first weight for the first node group, and
wherein a second measure of variation, that is lower than the first measure of variation, between the first node group and the second node group is associated with a third weight, that is higher than the first weight, for the first node group.
11 . The non-transitory computer-readable medium of claim 8 , wherein modifying the particular machine learning model based on the aggregated refinement information and the first and second weights respectively associated with the first and second node groups consumes fewer processing resources than modifying the particular machine learning model based on the plurality of instances of refinement information.
12 . The non-transitory computer-readable medium of claim 8 , wherein the refinement information includes feature importance gradients for one or more features associated with the particular machine learning model.
13 . The non-transitory computer-readable medium of claim 8 , wherein the measure of error associated with the second set of refinement information is a second measure of error, wherein determining the first weight associated with the first node group includes:
determining a first measure of error with respect to a particular label associated with a particular instance of refinement information associated with the first node group, wherein the first weight associated with the first node group is further based on the determined first measure of error.
14 . The non-transitory computer-readable medium of claim 8 , wherein the first set of nodes includes a plurality of Internet of Things (“IoT”) devices that each include one or more respective sensors that detect the sensor data associated with each respective node of the first set of nodes.
15 . A method, comprising:
receiving policy information specifying respective sets of criteria for a plurality of different node groups, wherein the respective sets of criteria include a device type criteria;
identifying that a first set of nodes is associated with a first device type;
identifying, based on a first device type criteria specified by the policy information and further based on identifying that the first set of nodes is associated with the first device type, that the first set of nodes are associated with a first node group of the plurality of node groups;
identifying that a second set of nodes is associated with a second device type;
identifying, based on a second device type of criteria specified by the policy information and further based on identifying that the second set of nodes is associated with the second device type, that the second set of nodes are associated with a second node group of the plurality of node groups;
receiving, from the first set of nodes, a first set of refinement information for a particular machine learning model, wherein the first set of refinement information includes a first plurality of modifications to a particular feature of the particular machine learning model, wherein each modification to the particular feature, of the first set of refinement information, has been generated by a respective node of the first set of nodes based on sensor data detected by the respective node of the first set of nodes;
receiving, from the second set of nodes, a second set of refinement information for the particular machine learning model, wherein the second set of refinement information includes a second plurality of modifications to the particular feature of the particular machine learning model, wherein each modification to the particular feature, of the second set of refinement information, has been generated by a respective node of the second set based on sensor data detected by the respective node of the second set of nodes;
generating first aggregated refinement information based on first set of refinement information for the particular machine learning model;
generating second aggregated refinement information based on the second set of refinement information for the particular machine learning model;
determining a first weight associated with the first node group, wherein the first weight is determined based on identifying that the first node group is associated with the first device type;
determining a second weight associated with the second node group, wherein the second weight is determined based on identifying that the second node group is associated with the second device type;
identifying a measure of error associated with the second set of refinement information, wherein identifying the measure of error includes:
applying the second set of refinement information to a set of data that includes a plurality of labels, and
identifying the measure of error based on applying the second set of refinement information to a particular subset, of the set of data, that is associated with a particular label of the plurality of labels;
identifying that the measure of error associated with the second set of refinement information exceeds a threshold measure of error;
determining a third weight associated with the second group based on identifying that the measure of error, associated with the second refinement information, exceeds the threshold measure of error; and
modifying the particular machine learning model based on the first and second aggregated refinement information and further based on the first and third weights.
16 . The method of claim 15 , wherein determining the first weight associated with the first node group includes:
determining a measure of variation between the refinement information associated with the first node group and the second node group,
wherein a first measure of variation between the first node group and the second node group is associated with the first weight for the first node group, and
wherein a second measure of variation, that is lower than the first measure of variation, between the first node group and the second node group is associated with a third weight, that is higher than the first weight, for the first node group.
17 . The method of claim 15 , wherein modifying the particular machine learning model based on the aggregated refinement information and the first and second weights respectively associated with the first and second node groups consumes fewer processing resources than modifying the particular machine learning model based on the plurality of instances of refinement information.
18 . The method of claim 15 , wherein the refinement information includes feature importance gradients for one or more features associated with the particular machine learning model.
19 . The method of claim 15 , wherein the measure of error associated with the second set of refinement information is a second measure of error, wherein determining the first weight associated with the first node group includes:
determining a first measure of error with respect to a particular label associated with a particular instance of refinement information associated with the first node group, wherein the first weight associated with the first node group is further based on the determined first measure of error.
20 . The method of claim 15 , wherein the first set of nodes includes a plurality of Internet of Things (“IoT”) devices that each include one or more respective sensors that detect the sensor data associated with each respective node of the first set of nodes.