Partitioning neural network training across devices using partitioning schedules
The system addresses the complexity of training large neural networks by using a partitioning schedule with both manual and automatic tactics to efficiently partition the training across multiple devices, enhancing training efficiency and adaptability.
Patent Information
- Application Number
- PCT/EP2024/085351
- Authority / Receiving Office
- WO · WO
- Patent Type
- Applications
- Current Assignee / Owner
- Priority Date
- 2023-12-07
- Filing Date
- 2024-12-09
- Publication Date
- 2025-06-12
AI Technical Summary
Existing techniques for training large neural networks on multiple hardware devices are complex and difficult to compose, often requiring manual specification and verification of partitioning strategies, which limits their efficiency and predictability.
A system that uses a partitioning schedule with a sequence of tactics to determine how to partition the training of a neural network across multiple devices, allowing for both manual and automatic tactics and enabling efficient and predictable composition of parallelization strategies.
The system simplifies the composition of partitioning strategies, provides performance estimates without execution, and allows for adaptive partitioning of the same model across different devices without modifying the underlying code, resulting in improved neural network training efficiency.
Smart Images

Figure EP2024085351_12062025_PF_FP_ABST
Abstract
Description
[0001]DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationPARTITIONING NEURAL NETWORK TRAINING ACROSS DEVICES USING PARTITIONING SCHEDULES CROSS-REFERENCE TO RELATED APPLICATION This application claims priority to GR Patent Application No.20230101009, filed on December 7, 2023. The disclosure of the prior application is considered part of and is incorporated by reference in the disclosure of this application. BACKGROUND This specification relates to training neural networks on multiple hardware devices. Neural networks are machine learning models that employ one or more layers of nonlinear units to predict an output for a received input. Some neural networks include one or more hidden layers in addition to an output layer. The output of each hidden layer is used as input to the next layer in the network, i.e., the next hidden layer or the output layer. Each layer of the network generates an output from a received input in accordance with current value inputs of a respective set of parameters. SUMMARY This specification describes a system implemented as computer programs on one or more computers in one or more locations that determines a partitioning of the training of a neural network across multiple hardware devices. Particular embodiments of the subject matter described in this specification can be implemented so as to realize one or more of the following advantages. Large neural networks are generally trained on multiple hardware devices through a combination of parallelization strategies. For example, these strategies can include two or more of data, model, or optimizer sharding. Each of these is a different way in which the training of a neural network may be partitioned across the multiple hardware devices. As such strategies become more complicated, existing techniques make it difficult to compose multiple techniques and often require users to manually specify each technique in the composition and verify that the resulting partitioning scheme is valid. This specification, on the other hand, describes an expressive and predictable partitioner that makes it simple to compose various strategies. In some cases, the partitioner can provide performance estimates for a given composition without execution of the underlying scheme. DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationIn particular, in the described specification, partitioning is driven by a schedule that specifies a sequence of partitioning tactics. Advantageously, these tactics are specified separately from the model code, allowing the described system to partition the same model differently on different devices without requiring rewriting the underlying code. By implementing the tactics through compiler rewrites, the described techniques effectively and efficiently implement complex compositions of partitioning strategies, resulting in improved neural network training. Moreover, the schedule can include both manual and automatic tactics, where the automatic tactics are automatically resolved by the system during the partitioning process. Thus, the system can benefit from both a predetermined, well-known sharding strategy manually inserted by a user and an automated one, e.g., resolved using a search, to discover unknown, higher-performing sharding strategies for a given neural network that are specific to a given set and configuration of devices. According to an aspect, there is provided a method performed by one or more computers, the method comprising: receiving mesh data specifying a configuration of a plurality of hardware devices that assigns each of the plurality of hardware devices to a respective index along each of one or more axes; obtaining model specification data characterizing a training step for training a neural network, the specification data identifying a plurality of tensors, the plurality of tensors comprising: (i) a plurality of tensors that are provided as input to the training step; and (ii) one or more tensors that are generated as output of the training step; obtaining a schedule for partitioning the training of a neural network across the plurality of hardware devices, wherein the schedule comprises a sequence of partitioning tactics, and wherein each partitioning tactic defines, for each of one or more tensors of the plurality of tensors, a partitioning of a respective dimension of the tensor across one of the one or more axes; and processing the mesh data, the model specification data, and the schedule to generate program data that, when executed by each of the plurality of hardware devices, partitions the training step across the plurality of hardware devices. In some embodiments, the method further comprises providing the program data to a compiler for compilation to generate compiled program data. In some embodiments, the method further comprises causing the plurality of hardware devices to execute the compiled program data generated by the compiler in order to perform one or more training steps for training the neural network. Each of the hardware devices may be associated with one of the indices. DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationThe details of one or more embodiments of the subject matter of this specification are set forth in the accompanying drawings and the description below. Other features, aspects, and advantages of the subject matter will become apparent from the description, the drawings, and the claims. BRIEF DESCRIPTION OF THE DRAWINGS FIG.1 is a diagram of an example partitioning system. FIG.2 is a flow diagram of an example process for portioning the training of a neural network. FIG.3 shows an example of batch and model parallelism. FIG.4 shows an example of a simplified view of a training step with batch and optimizer parallelism. FIG.5 shows an example of an architectural overview of the partitioning system. FIG.6 shows an example of the semantics of all_slice and all_gather via a chain of operations. Like reference numbers and designations in the various drawings indicate like elements. DETAILED DESCRIPTION FIG.1 is a diagram of an example partitioning system 100. The partitioning system 100 is an example of a system implemented as computer programs on one or more computers in one or more locations, in which the systems, components, and techniques described below can be implemented. The system 100 determines a partitioning of the training of a neural network across multiple hardware devices 130. Generally, the multiple hardware devices 130 can include any appropriate devices that can carry out operations required to perform neural network training. Examples of such devices include central processing units (CPUs) and hardware devices that include one or more hardware accelerators that have circuitry for performing matrix-vector multiplication in hardware, e.g., graphics processing units (GPUs), tensor processing units (TPUs), and other ASICs that are optimized for performing machine learning computations. In particular, the system 100 receives mesh data 140 that specifies a configuration of the plurality of hardware devices 130. DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationThe configuration assigns each of the plurality of hardware devices 130 to a respective index along each of one or more axes of a mesh. For example, the mesh can be a one-dimensional mesh, so that each device is assigned an index along a single axis. As another example, the mesh can be a two-dimensional mesh, so that each device is assigned a respective index along each of the two axes. As yet another example, the mesh can be a three-dimensional mesh, so that each device is assigned a respective index along each of the three axes. Generally, the configuration specified by the mesh is a logical configuration of the hardware devices but can be laid out to reflect the underlying system’s communication topology. In more detail, the distributed execution of large tensor programs, e.g., programs that operate on tensors to train neural networks, employs the concept of a “mesh.” A mesh is an n-dimensional array with named axes that offers a logical view of the available devices. For example, a system of 16 devices can be represented as a 2D mesh with axes a and b{^^:2,^^:8}; or as 3D mesh with axes a, b, and c {^^:2,^^:2,^^:4}; or just as a 1D mesh with a single axis a {^^:16} among many others. Generally, the mesh is laid out respecting the underlying system’s communication topology to make reasoning about performance easier and to exploit the relative speeds of the networks that connect the devices better. In other words, the logical arrangement of devices as specified by the indices of the mesh aligns with the physical connectivity of the devices. For example, consider a cluster of 4 servers with 8 accelerator devices each, connected through a fast interconnect network within a server but slower Ethernet connections across servers: one may want to view this system as a ^^ :4,^^:8 system so that communication along devices that span each axis of the mesh happens exclusively on one of the two networks. These numbers are examples only. More generally, the mesh may be arranged with a first axis, wherein a first number of devices is arranged along the first axis, and a second axis, wherein a second number of devices is arranged the second axis. The devices may communicate with other ones of the devices arranged along the first axis according to a first communication link type and may communicate with other ones of the devices arranged along the second axis according to a second communication link type. The system 100 also obtains model specification data 150 characterizing a training step for training the neural network. A “training step” as used in this specification refers to DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT Applicationthe processing required to train the neural network on a batch of training data and generally involves updating the parameters, e.g., weights, of the neural network by performing a forward step followed by a backward step through the neural network. The overall training process performed over multiple batches of training data and one or more epochs may comprise a plurality of these training steps. Generally, the specification data 150 identifies a plurality of tensors that are processed by performing the training step and that are generated as a result of performing the training step. The tensors generally include a plurality of tensors that are provided as input to the training step, e.g., tensors representing the weights of the neural network, the training inputs (e.g. training data provided as input to the neural network) for the training step, and any other quantities that are maintained across training steps. For example, these other quantities can include quantities required to maintain the state of the optimizer that is being used to train the neural network. Each of at least some of the tensors may correspond to a variable in the specification data. The tensors also generally include one or more tensors that are generated as output of the training step, e.g., tensors representing the updated weights of the neural network and any other quantities that are maintained across training steps. For example, these other quantities can include quantities required to maintain the updated state of the optimizer that is being used to train the neural network after the training step is performed. The state of the optimizer may be otherwise referred to as optimizer state (e.g. ADAM or Adafactor optimizer state) and can include, e.g., one or more tensors specifying one or more moments of the gradients with respect to the network parameters. For example, the specification data 150 can be generated by a machine learning framework from a representation of the neural network generated in the machine learning framework, e.g., a representation of the neural network as a computational graph, or provided as input to the machine learning framework. Examples of machine learning frameworks that can be used to provide an input to the system 100 include JAX and Tensorflow. The system 100 obtains a schedule 160 for partitioning the training of the neural network across the plurality of hardware devices. As will be described below, in some cases, the schedule 160 is generated by the system 100 while, in other cases, the schedule 160 is generated by the machine learning framework or by a user and is provided as input to the system 100. DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationAs a particular example, the system 100 can expose an application programming interface (API) through which a user or another system can submit the schedule 160, e.g., as an ordered list of partitioning tactics. The schedule 160 includes a sequence of partitioning tactics. Each partitioning tactic defines a partitioning of at least a respective dimension of a respective one of the plurality of tensors specified in the specification data 150 across one of the one or more axes of the mesh. In some cases, a given partitioning tactic can define a partitioning of respective dimensions of two or more tensors across a given axis. The partitioning tactics can include only manual partitioning tactics, only automatic partitioning tactics, or both manual and automatic partitioning tactics. A manual partitioning tactic is one that identifies both (i) the dimension to be partitioned and (ii) the mesh axis across which to partition the dimension. The dimension and mesh axis are specified for each manual partitioning tactic by user input. An automatic partitioning tactic is one that only identifies a mesh axis, but does not identify which dimension of which tensor to partition across the dimension. For example, as described below, the mesh axis may be specified by user input, but the dimension of the tensor is determined automatically by the system based on a runtime cost model. Manual and automatic partitioning tactics will be described in more detail below. Advantageously, the schedule 160 can be generated independently from the model code that is processed by the machine learning framework and can be adapted for the particular configuration of devices 130 without modifying the model code. That is, the machine learning framework can maintain a single representation of the model and only the schedule 160 needs to be modified in order to perform the training of the neural network on a different set of available devices, i.e., in order to optimize the partitioning of the training of the neural network for a new set of devices. The system 100 processes the mesh data, the model specification data, and the schedule to generate program data 170 that, when executed by each of the plurality of hardware devices, partitions the training step across the plurality of hardware devices. For example, when the system 100 operates under a single program multiple device (SPMD) model, the program data 170 can specify a program to be executed by each of the devices to perform the training step. Processing the mesh data, the model specification data, and the schedule to generate program data is described in more detail below. DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationThe system 100 can then cause the program data 170 to be executed by the hardware devices 130 in order to cause the hardware devices 130 to perform the training step for training the neural network. In some cases, the program data 170 can be static, i.e., the same program data can be used for each training step in training the neural network, and the execution of the program data causes the hardware devices 130 to perform a sequence of multiple training steps. For example, the system 100 can provide the program data 170 to a compiler for compilation. Examples of suitable compilers that can be used to compile the program data 170 for execution include XLA and OpenXLA. After compilation, the system 100 can cause the plurality of hardware devices 130 to execute compiled program data generated by the compiler in order to perform one or more training steps for training the neural network. FIG.2 is a flow diagram of an example process 200 for partitioning the training of a neural network across devices. For convenience, the process 200 will be described as being performed by a system of one or more computers located in one or more locations. For example, a partitioning system, e.g., the partitioning system 100 depicted in FIG.1, appropriately programmed in accordance with this specification, can perform the process 200. The system receives mesh data specifying a configuration of a plurality of hardware devices that assigns each of the plurality of hardware devices to a respective index along each of one or more axes (step 202). The system obtains model specification data characterizing a training step for training a neural network (step 204). As described above, the specification data identifies a plurality of tensors that include (i) a plurality of tensors that are provided as input to the training step and (ii) one or more tensors that are generated as output of the training step. The system obtains a schedule for partitioning the training of a neural network across the plurality of hardware devices (step 206). The schedule includes a sequence of partitioning tactics. Each partitioning tactic defines, for each of one or more tensors of the plurality of tensors, a partitioning of a respective dimension of the tensor across one of the one or more axes. The schedule of tactics can include one or more manual partitioning tactics. Each manual partitioning tactic identifies a respective axis and identifies, for each of one or more DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT Applicationtensors of the plurality of tensors, the respective dimension of the tensor to be partitioned across the axis. The schedule of tactics can also include one or more automatic partitioning tactics. Each automatic partitioning tactic identifies a respective axis but does not identify the respective dimension of the tensor to be partitioned across the respective axis. In some cases, the schedule of tactics includes both manual and automatic partitioning tactics. The system processes the mesh data, the model specification data, and the schedule to generate program data that, when executed by each of the plurality of hardware devices, partitions the training step across the plurality of hardware devices (step 208). The program data may specify the partitioning of the training step across the plurality of hardware devices such that the operations involved in the training step are accordingly partitioned across the plurality of hardware devices when the program data is executed. The partitioning of the training step determines which elements of the input tensors are processed on which hardware device and which elements of the output tensors are produced by which hardware device. Generally, as part of this processing, the system maps each partitioning tactic in the schedule to a sequence of compiler actions. The system can then propagate the compiler actions through the program data based on linear algebra homomorphisms defined by the partitioning tactics in the schedule. In particular, the system can perform a propagation pass through the program data that creates loops around operations and, optionally, slices other arguments within the program data. The system can maintain a registry or other data that encodes linear algebra homomorphisms. For example, the registry can encode a homomorphism as a map between tiling and reduction attributes. The system can then use the encoded homomorphisms to modify the program data, e.g., by introducing more loops, during the propagation pass, improving the effectiveness of the implementation of the partitioning tactics. For example, homomorphisms can allow an operation to be rewritten in order to introduce additional tiling in accordance with a given partitioning strategy, improving parallelization when executed on the devices. Performing a propagation pass and the loop and slice operations that are involved are described in more detail below. Making use of linear algebra homomorphisms is also described in more detail below. DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationAs part of this, the system can incrementally resolve conflicts caused by partitioning tactics according to the order of the partitioning tactics within the sequence. That is, the system can perform rewriting of the program data incrementally, so that if a conflict is encountered, the system will refrain from implementing the partitioning tactic that appears later within the sequence while still implementing the partitioning tactic that appears earlier. Resolving conflicts is described in more detail below. As described above, automatic partitioning tactics do not identify which dimension of which tensor is to be partitioned across the axis identified by the partitioning tactic. Therefore, when the schedule includes an automatic partition tactic as part of the processing, the system performs a search to identify a particular dimension of a particular tensor to be partitioned across the axis identified by the automatic partitioning tactic. For example, the system can perform the search using a runtime cost model. That is, the system automatically determines, by using the cost model, the partitioning that would be most beneficial in terms of minimizing runtime cost. Examples of runtime cost models that can be used by the system to perform the search are provided in: Sami Alabed, Dominik Grewe, Juliana Franco, Bart Chrzaszcz, Tom Natan, Tamara Norman, Norman A. Rink, Dimitrios Vytiniotis, and Michael Schaarschmidt. Automatic discovery of composite spmd partitioning strategies in partir, 2022 and Michael Schaarschmidt, Dominik Grewe, Dimitrios Vytiniotis, Adam Paszke, Georg Stefan Schmid, Tamara Norman, James Molloy, Jonathan Godwin, Norman Alexander Rink, Vinod Nair, et al. Automap: Towards ergonomic automated parallelism for ml models, arXiv preprint arXiv:2112.02958, 202. For example, such a runtime cost model may estimate the peak memory usage, runtime duration and communication resources consumed for the partitioning of different dimensions of a particular tensor. On this basis, the system may select one of the dimensions of the tensor to partition. Thus, this allows users to provide the system with known partitioning strategies that are likely to improve the execution of the model training for some of the tensors while allowing the system to search for partitioning strategies for certain mesh dimensions to further optimize the training. The system can then cause the program data to be executed by the hardware devices in order to cause the hardware devices to perform the training step for training the neural network. In some cases, the program data can be static, i.e., the same program data can be used for each training step in training the neural network, and the execution of the program data causes the hardware devices to perform a sequence of training steps. DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationFor example, the system can provide the program data to a compiler for compilation. After compilation, the system can cause the plurality of hardware devices to execute compiled program data generated by the compiler in order to perform one or more training steps for training the neural network. In some implementations, when performing the training step, the plurality of hardware devices operate under a single program multiple data (SPMD) model and each execute a same program. In these implementations, the program data specifies the same program to be executed by each of the plurality of devices to perform the training step. That is, the program data specifies a single program and, when performing a given training step, the different devices execute the single program, but on different data. For example, when the devices operate under the SPMD model, the program data can include device-local SPMD code for execution by each of the plurality of hardware devices. For example, generating device-local SPMD code can include “lowering” higher-level representations to lower-level representations. Examples of performing the lowering are described below. By generating program data as described above, the system can effectively implement a mixture of any of a variety of parallelism strategies for more effectively partitioning the training of the neural network. Some examples of partitioning strategies will now be described. One example of a partitioning strategy is batch parallelism. In batch parallelism (BP), the input batch is sharded across the devices while the model parameters and optimizer state are replicated. This then requires one AllReduce operation per parameter in the backward pass. Another example of a strategy is model parallelism (MP). One form of model parallelism occurs when the parameters of the neural network are sharded across the devices. FIG.3 shows an example 300 of batch and model parallelism. In particular, FIG.3 shows an example of M-way model parallelism, i.e., M-way parameter sharding, and N-way batch parallelism, i.e., N-way data sharding. On the left, the communication along the M axis (e.g., activation reductions) is shown. On the right, the communication along the N axis (e.g., gradient reductions) is shown. Notably, all devices along the batch axis hold the same shard of the model parameters. Another example of a strategy is optimizer sharding. In optimizer sharding, the optimizer state is sharded, e.g., to reduce the peak memory use of the training process. In one variant of optimizer sharding, the parameter gradients are also sharded across devices. In DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT Applicationanother variant, both the parameter gradients and the parameters are sharded across devices. For example, an SPMD-oriented variant of this shards all parameters, resulting in an approach known as “fully-sharded data parallelism” (FSDP). From a communication perspective, this introduces an AllGather operation whenever the parameters are needed, i.e., once in the forward and once the in backward pass for every parameter, but has the advantage that the full parameters must live on the device briefly, just before they are used. Correspondingly, in the backward pass, the gradients for the parameters are not reduced; they are rather reduce-scattered across the devices – a cheaper operation. The optimizer is then also sharded and uses the gradient shards and parameter / optimizer shards to update the parameter shards. FIG.4 shows an example 400 of a simplified view of a training step with batch and optimizer parallelism. In particular, the top of the FIG.4 shows batch parallelism while the bottom shows optimizer / FSDP parallelism, assuming two devices. In the bottom, the parameters are all-gathered before their use, and parameter gradients are reduced-scattered before being used to update the local parameter shards. The aforementioned partitioning strategies can be sequentially combined to span one or multiple axes of the device mesh. For example, FIG.3, described above, illustrates partitioning over a 2D mesh with data parallelism over one axis and model parallelism over the other. Generally, for training large models, multiple strategies are composed. For example, a system can compose batch parallelism, model parallelism, and FSDP to generate an overall strategy for training a large neural network on a given set of devices. Beyond the strategies described above, there exist many other strategies for sharding one or more tensors during the training of a large neural network. Other examples include activation sharding after model parallelism, multi-query sharding in Transformers, and Transformer sequence sharding. By making use of the techniques described in this specification, the system can effectively and automatically determine combinations of these strategies in ways that improve the effectiveness of the training of the neural network. Additionally, the system can verify that each strategy was applied correctly, as the user can inspect the collectives introduced and estimated performance after each strategy. More specifically, partitioning strategies can be represented as program transforms, i.e., modifications to the program executed by the devices to perform the training steps. DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationAs an example, consider a JAX program of two matrix multiplications: The jax.jit() call traces the Python function to construct a pure representation that will later be converted to an intermediate representation used by the compiler, e.g., an XLA HLO (High Level Optimizer) representation, get compiled for a specific backend, and run. While the remaining A simplified version of the program that would result from tracing the original Python JAX function is as follows: In this representation, each value is annotated with a shape, e.g., % x1 = ...: tensor<256x8xf32>, denoting an array of shape [256,8] and element type of float32. This program can then be run on a mesh of devices based on a parallelization strategy. As an example, assume the program is to be executed on a 2D mesh { : 4,^^ : 2}. For batch parallelism, one strategy is to partition the first (256-sized) dimension of input %x across the mesh axis ^^, observing that the whole program is a pure map over that dimension: The resulting device-local program above takes a first argument of smaller shape 64x8 since each device acts in parallel on a slice of %x determined by the device coordinate along axis ^^. At the same time, the shape of parameters %w1 and %w2 remains constant across all devices (due to replication) – in effect, this transformation is batch parallelism. DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationThe program can be further modified to add model parallelism. For example, if, on top of BP, the parameter %w1 is partitioned on ^^^^^^ = 1 and %w2 on ^^^^^^ = 0 along axis ^^, then each device along axis ^^ may perform a smaller multiplication with a different parameter shard: The first matmul is a map over the second dimension of %w1, so no special care is necessary. However, note that %x1 is partitioned along its second dimension. For the second dimension, both of its operands are partitioned along the contracting dimensions. Hence, the original program semantics can be recovered by inserting a final all_reduce operation across axis ^^. Note, however, that the parameters are only sharded on axis ^^ (but not ^^) in the above program. The function can be modified further to shard the parameters further on dimensions 0 and 1, respectively, but this would require inserting two all_gather operations, right before they are needed in the matrix multiplication: These operations gather the shards on the corresponding dimensions. This sharding of parameters, after batch parallelism, yields the optimizer sharding described above. Note that the parameters are only gathered before use, reducing the peak memory requirements. As described above, more complex shardings are possible on top of the aforementioned chain of strategies. For example, the input and output activation (%x and the return value) may be additionally sharded on the model axis ^^. Such a sharding will convert DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT Applicationthe all_reduce to a all_reduce_scatter and introduce an all_gather on the input %x before using it in the first matmul. The sequence of examples above highlights that (i) popular sharding strategies do compose, (ii) by following linear algebraic reasoning, a model can be partitioned just by sharding its inputs and parameters, and optionally an internal operation (iii) it suffices to consider simple properties about the parallel and the contracting dimensions of the operations in the program. As described above, by making use of the described techniques, these strategies can be expressed as compiler actions generated by the system that enable a semantics-preserving rewrite of an initial unpartitioned program, e.g., Listing 2 shown above, into a device-local program, e.g., Listings 3 to 5. That is, by operating on a schedule of partitioning tactics, the described techniques can transform the initial unpartitioned, e.g., Listing 2, into any of a variety of device-local programs that are optimized for different devices without needing to modify the underlying model code. While the concepts are described above with reference to a simple sequence of matrix multiplications, a typical machine learning model contains thousands of tensor operations (e.g., dot-products, convolutions, gather and scatterops, and more) and control flow (e.g., and loops). A full training step, including back-propagation and optimizer, can reach 10-100k operations and have thousands of multi-dimensional array arguments for model parameters and optimizer state. As a result, the described techniques allow users to effectively partition these operations and arrays without requiring the user to manually specify the desired partitioning for each. In particular, the system can expose an API or other interface that allows users to compose partitioning strategies incrementally. In particular, as described above, the system obtains schedule data that specifies a schedule, which is a sequence of manual or automatic or both tactics. DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationFor example, to achieve the final sharding of Listing 1 that is described in Listing 5, a JAX user can submit the following through an interface exposed by the system: That is, the mesh data is described by the maps.mesh function while the schedule is the sequence of tactics BP, MP, and Z3. The first tactic BP partitions the first argument ’x’ on the 0 dimension and across axis "B"; yielding batch parallelism (Listing 3). The second tactic MP partitions input w1 on dimension 1. As result, the compiler will determine that this action is sharding the contracting dimension of a matmul and will also end up sharding w2 on dimension 0 through a process referred to as propagation invoked at the end of every tactic. The final tactic Z3 shards the parameters on the remaining available dimensions and axis ^^ (Listing 5). As can be seen above, each of these three tactics is a manual tactic that identifies both the tensor dimension(s) and the mesh axis. As will be described below, a given schedule can also include automatic partitioning tactics. The schedule is then passed to the system together with the function f to partition and the device mesh (the output of the maps.mesh function). The system can then trace the function into an intermediate representation module that the system then transforms according to the schedule. As a particular example, the system can return a partitioned module exposed as a Python callable ready to be called with sharded arrays and executed on the devices of the mesh. Additionally, the system produces various metadata that includes a sharding specification of the function inputs and outputs produce. Optionally, the system can also output cost estimates (e.g., collective communication breakdown by type and simulation results) recorded after every tactic in the schedule. DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationAs described above, in addition to manual partitioning tactics, the system exposes, i.e., allows users to specify, an AutomaticPartition() tactic over one or more mesh axes, which can be composed with other manual tactics. For example, users could submit a schedule as follows: This is an example of a schedule with both manual and automatic scheduling tactics. The first tactic partitions a module, introducing batch parallelism manually, but the second tactic causes the system to use a search algorithm to discover an optimal partitioning along axis "M.” For example, the system can perform this search using a runtime cost model that attempts to identify the lowest-cost model, e.g., in terms of peak memory use, latency, or both, while penalizing models that do not fit, i.e., penalizing partitionings that result in the operations assigned to any given device exceeding the amount of memory available to the device or otherwise violating a constraint on execution, effectively automating the discovery of optimal model parallelism. To ensure composability, both manual and automatic tactics issue sequences of (the same) lower-level compiler actions that either (i) shard a value dimension along an axis or (ii) explicitly keep a value unpartitionable across a mesh axis, or (iii) propagate sharding information across a module. For example, the example schedule that includes the sequence of tactics BP, MP, and Z3above generates a sequence of seven actions at the compiler level: FIG.5 shows an example 500 of an architectural overview of the partitioning system. The components belonging to example 500 may be software components provided by the partitioning system 100. In the example 500, the system receives data from a machine learning framework and provides as output, program data to be compiled by the compiler by DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT Applicationthe partitioning system 100. In this example, the machine learning framework is described as a JAX machine learning framework and the compiler is described as being an XLA compiler. The example 500 describes an implementation of the partitioning in the Multi-Level intermediate representation (IR) Compiler (MLIR) framework. Accordingly, in the example 500, "+" is used to signify the introduction of new operators, and "-" to signify that certain operators have now become illegal. Moreover, components of the system that are introduced on top of the MLIR framework as referred to in the example 500 as “PartIR” components, e.g., PartIR:Core, PartIR:MHLO, PartIR:SPMD, ParIR:Temporal, and so on. Thus, in the example of FIG.5, programs are generated from tracing JAX functions into an intermediate representation dialect used by the system. In the example 500, this is the MLIR StableHLO dialect. Manual or automatic tactics from the received schedule generate sequences of compiler actions that introduce and propagate functional StableHLO loops and specialized slicing ops, that belong in the PartIR:Core dialect. PartIR:Core loops and slices are interpreted as sequential loops in the PartIR:Temporal dialect, whose main use is a reference semantics of PartIR:Core alongside optional applications like automatic microbatching transforms. PartIR:Core loops and slices are canonically lowered to the PartIR:MHLO dialect to generate device-local collective communication ops. During the lowering process, PartIR:SPMD organizes the loops to capture the distribution and partial replication of values in a global view, allowing the system to eliminate any needless redistribution of tensors. The collectives at PartIR:MHLO refers to mesh axes that make their IR encoding independent of the total number of devices in the mesh (as opposed to collectives in StableHLO and XLA HLO that directly reference groups of logical device IDs), and makes it easy to reason about and fuse. Optionally, the system can perform simulation and cost estimation at this level. To export the final program data, the system lowers any custom high-level PartIR:MHLO ops to StableHLO computations and hand over the module to XLA for compilation. As the figure shows, the architecture and separation of dialects employed by the system allows the system to implement the right rewrites at the right level of abstraction and to also independently test various components of the partitioner. For example, PartIR:Core is DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT Applicationunaware of SPMD execution, and its rewrite axioms remain very simple; this is a separate step from SPMD lowering and optimization of communication. PartIR:Core PartIR:Core introduces just two operations on top of the legal operations in StableHLO: (1) a loop operation that expresses functional parallel tiling or reduction loops, and (2) a slice operation that extracts a tensor slice along some dimension based on a loop index. A tiling action tile<%value, dim, axis> is reified in the intermediate representation by creating a loop that returns, in each iteration, a slice of %value along dimension dim. For example, value tiling %x along dimension 0 and axis "B" from earlier produces: The loop operation contains two static attributes: (i) a mesh axis ("B") and (ii) an action attribute (#tile<0>). It also accepts a single-argument closure (region in the MLIR jargon) that represents the loop body: (%rB: range<4>){... }. The closure takes as an argument a range value (%rB) and performs a tensor computation returning a value of type tensor<64x8xf32>. The range argument %rB plays the role of a loop index given an PartIR- specific range type. slice operations consume these loop indices. The meaning of slice 0 %x[%rB] is that it extracts the %rB-th chunk of the tensor %x along dimension 0. The tiling here perfectly partitions the tensor dimension into 4 equally-sized, contiguous chunks since axis "B" has size 4. Hence, the result of slice has shape 64x8. Furthermore, the value tiling action has replaced %x of type tensor<256x8xf32> with value %xt of the same type – i.e., value tiling is a semantics- and type-preserving local rewrite. Note that there is nothing specific to SPMD execution about PartIR:Core constructs, and thus far, tiling loops can be considered just concatenating the chunks from each iteration. Value tiling actions simply create loops around sliced values, so they only reify a tiling action in the IR. However, they help bootstrap a powerful propagation pass that consequently creates loops around operations and further slices other arguments. This propagation is justified by program equivalences that directly encode linear algebra homomorphisms. DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationPartIR propagation is built around program equivalences that involve loop and slice instructions. Three admissible program equivalences for a matrix multiplication are shown below: The first two rewrite a matrix multiplication as a tiling loop with a smaller multiplication inside. The last one introduces a new form of attribute accompanying the loop, a #sum action attribute. This signifies that the results of each iteration of the loop should be reduced, as the system is slicing the operands on their contracting dimension. These equivalences are justified as linear algebra homomorphisms, with the monoidal structure stacking for the first two and addition for the last one. To allow the system to implement the rewriting code, for all operators, the system equips PartIR with a tile-mapping registry (TMR) that enables a concise encoding of linear algebra homomorphisms as maps between tiling and reduction attributes. The TMR contains, for every tensor operation with ^^ inputs, a set of specifications of the form: ^^⊥ 1 , ... , ^^⊥ ^^ ↩→^^1, ... ,^^ where ^^⊥ stands for an optional tiling action, while ^^ stands for an arbitrary action (including #sum). StableHLO ops and the described loops may return multiple results, hence the generalized form ^^1, ... ,^^ . One such specification asserts that a given operation can be rewritten as a loop with action(s) ^^1, ... ,^^ if the system slices its operands according to ^^⊥ 1 , ... , ^^⊥ ^^ (a missing action implies no slicing). DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationBelow are entries for matmul (corresponding to the three equivalences from Figure 4) and for an elementwise add operation that asserts that tiling its result requires tiling its operands in the same way: TMR(matmul) = {(#tile^0^,⊥) ↩→#tile^0^} ∪ {(⊥,#tile^1^) ↩→#tile^1^} ∪ {(#tile^1^,#tile^0^) ↩→#sum} TMR(add) This abstraction is sufficient to capture a wide variety of equivalences, including substantially more complex operations such as convolutions, generalized dot products, scatter and gather operations, dynamic slicing, and reshapes. Addition is assumed here, but the implementation supports #sum<@f> custom reductions for any associative reduction @f. From an implementation standpoint, since the semantics of many MHLOperations are determined by static attributes and sometimes the shapes of their operands, the TMR implementation keys on the operation, its static attributes, and operand shapes, and sometimes the TMR also registers how attributes should additionally be transformed. Propagation is a pass that greedily attempts to propagate known and partially known information and introduce more loops-based on the linear algebra homomorphisms in the TMR. Forward propagation searches for an entry that matches the actions of loops that produce operands of an operation, whereas backward propagation searches for an entry that matches the way the operation result is sliced downstream. For example, assume that the value %x1 in Listing 2 has been tiled through a value tiling action. Propagation is a pass that greedily attempts to propagate known and partially known information and introduce more loops-based on the linear algebra homomorphisms in the TMR. DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationTo propagate tiling forward, observe that the matmul defining %x2 now takes an operand whose first dimension is tiled, hence the TMR entry (#tile^0^,⊥) ↩→#tile^0^ for matmul matches its operand context. Propagating backward, %x1 is produced by a matmul and is then sliced along dimension 0, which matches the result of that same TMR entry. Through propagation, the system has arrived at a program where every operation is within a loop context. These loops may be fused (important for temporal lowering) or directly lowered to SPMD. Here is what the fused program looks like: Optionally, as part of generating the program data, the system can perform a process referred to as inference. Inference is the process of deducing missing operand value tiling based on a partial match against a TMR entry. Continuing from the above fused program consider value-tiling %w2: The TMR entry #tile^1^,#tile^0^) ↩→#sum is a partial match on the operands of the second matmul, since the second operand (%w2t) is already tiled. This can be extended into a full match by value tiling the first operand (%x1s) and then continuing with forward propagation to arrive at: DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT Application In that final program, both %w1 and %w2 end up only used sliced across axis "M", even though only %w2 was explicitly value tiled. Inference is of practical importance in ML programs. For example, the optimizer state in a training step function can be left to be partitioned as the parameters are by inference since parameters and optimizer state flow to element-wise operations that express the parameter update. In some situations, it is impossible to propagate. One of these is when a loop over some mesh axis would need to be inserted inside an already inserted loop over the same axis, in which the system forbids maintaining a correspondence between loop nests and device meshes. Another one is when multiple (partial) TMR matches are found — a situation that is referred to as a conflict. Consider: Here, the operands of the matmul defining %x1 are tiled in a way that matches two TMR entries: (#tile^0^,⊥) ↩→ #tile^0^ and (⊥,#tile^1^) ↩→#tile^1^. The system will not attempt to resolve conflicts automatically. DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationInstead, the system can perform the rewriting incrementally. For example, performing value tiling on %x and propagating that choice before value tiling %w1 would yield the following program: At this point the TMR entry matches the definition of %x1s again. Alas, the operation in hand is already nested inside a loop over axis "B" and no further propagation is possible – creating a doubly-nested loop over "B" is invalid. This prioritization of BP over subsequent parameter sharding is exactly what is needed for, e.g., a ZeRO sharding strategies. The prioritization of rewrites, happening naturally at the boundaries of manual tactics lead to a rare occurrence of conflicts and as a result, dramatically reduces the need for many internal sharding decisions. As described above, PartIR:MHLO offers a device-local view of the computation. It permits the standard tensor and control flow operations of StableHLO / MHLO, but includes specialized operations for collective communication. These, unlike their low-level MHLO counterparts, operate on mesh axes, meaning that the communication spans across device groups defined by a different set of coordinates along these axes. Examples of these operations are as follows: These examples can be explained as follows: DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT Application• An all_reduce (Example 1) primitive reduces chunks along one or more mesh axes according to some reduction function (in the example @red_fn) and initial value. The return type is the same as the operand type. • An all_slice primitive (Examples 2 / 3) includes an array of axes per dimension, specifying the axes to be used to slice each dimension of the operand. In Example 2, the first dimension of size 16 is sliced by axis "x1", the second dimension is not sliced, and the third dimension is sliced by "x2". The result tensor type will have dimension sizes where each dimension is divided by the size of (the product of) the slicing axes in this dimension. Note that (Example 3) there can be multiple axes slicing the same dimension. The semantics of the op is that each device receives a slice of the original array based on the coordinate tuple along the slicing axes. • An all_gather primitive (Example 4) is the dual to all_slice – namely an array is gathered along the gathering axes in each dimension. Each dimension of the returned value will have size multiplied by (the product of) the gathering axes in this dimension. Figure 5 demonstrates the semantics of all_slice and all_gather via a chain of operations. • An all_reduce_scatter (Example 5) represents the fusion of an all_reduce and a subsequent all_slice. In particular, Example 5 is the fusion of Examples 1 and 2. • Finally, an all_to_all (Example 6) corresponds to a fusion of an all_gatheralong some dimension (0 in Example 6) followed by all_slice over the same sequence of axes in another dimension (1 in Example 6). FIG.6 shows an example 600 of the semantics of all_slice and all_gather via a chain of operations. In particular, the example 600 is a demonstration of sequences of all_slice and all_gather collectives on a mesh {x:2, y:2}. The top of FIG.6 shows that all devices hold the same 2D array, In the bottom, data is sliced row-wise along axis "y". On the right, data is further sliced column-wise along axis "x". In each case the the device-local tensor types are shown. The techniques described above can be accompanied with an aggressive fusion that aggressively fuses collectives together when possible. As described above, the system performs a lowering process, during which the system (PartIR:SPMD described above) neatly organizes the loops to capture the distribution and partial replication of values in a global view, allowing the system to eliminate any needless redistribution of tensors. More specifically, lowering of PartIR:Core to PartIR:MHLO is a typepreserving transformation that transforms (i) slice ops to all_slice collectives, and (ii) inserts all_gather DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT Applicationor all_reduce collectives in the results of tiling loops. This can be demonstrated with an example: Mechanically following the lowering rules gives: As a final step, PartIR:MHLO will aggressively optimize collectives; in this example all_gather and all_slice collectives fuse away. In addition the function arguments (and DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT Applicationfunction type) can be modified when the arguments are used sliced and the results are gathered: As described above, the system can optionally perform simulation. SPMD makes the collective explicit in relation to the mesh, with PartIR:MHLO providing the tensors that lives in each local device. This makes the simulator employed by the system simple: it iterates over the tensors in device-local context to estimate the performance metrics, with the aid of a hardware specification. Having the mesh axes encoded in the IR simplifies implementing optimizations that needs to reason with device groups by implementing them as one or more additional tactics. The user can control which axis to apply the optimization on, as well as any hyperparameters to tune. This additional tactic can be debugged as described above for the remaining tactics in the schedule. One example of exposing optimization as tactics follows: . This specification uses the term “configured” in connection with systems and computer program components. For a system of one or more computers to be configured to perform particular operations or actions means that the system has installed on it software, firmware, hardware, or a combination of them that in operation cause the system to perform the operations or actions. For one or more computer programs to be configured to perform particular operations or actions means that the one or more programs include instructions that, when executed by data processing apparatus (i.e. by at least one processor of the data processing apparatus), cause the apparatus to perform the operations or actions. Some of the above descriptions have been given in the context of particular frameworks (e.g. JAX or XLA), programming languages (e.g. Python). However, the DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT Applicationdescribed techniques are not limited to these frameworks, languages or particular standard protocols. Embodiments of the subject matter and the functional operations described in this specification can be implemented in digital electronic circuitry, in tangibly-embodied computer software or firmware, in computer hardware, including the structures disclosed in this specification and their structural equivalents, or in combinations of one or more of them. Embodiments of the subject matter described in this specification can be implemented as one or more computer programs, e.g., one or more modules of computer program instructions encoded on a tangible non transitory storage medium for execution by, or to control the operation of, data processing apparatus. The computer storage medium can be a machine- readable storage device, a machine-readable storage substrate, a random or serial access memory device, or a combination of one or more of them. Alternatively or in addition, the program instructions can be encoded on an artificially generated propagated signal, e.g., a machine-generated electrical, optical, or electromagnetic signal, that is generated to encode information for transmission to suitable receiver apparatus for execution by a data processing apparatus. The term “data processing apparatus” refers to data processing hardware and encompasses all kinds of apparatus, devices, and machines for processing data, including by way of example a programmable processor, a computer, or multiple processors or computers. The apparatus can also be, or further include, special purpose logic circuitry, e.g., an FPGA (field programmable gate array) or an ASIC (application specific integrated circuit). The apparatus can optionally include, in addition to hardware, code that creates an execution environment for computer programs, e.g., code that constitutes processor firmware, a protocol stack, a database management system, an operating system, or a combination of one or more of them. A computer program, which may also be referred to or described as a program, software, a software application, an app, a module, a software module, a script, or code, can be written in any form of programming language, including compiled or interpreted languages, or declarative or procedural languages; and it can be deployed in any form, including as a stand alone program or as a module, component, subroutine, or other unit suitable for use in a computing environment. A program may, but need not, correspond to a file in a file system. A program can be stored in a portion of a file that holds other programs or data, e.g., one or more scripts stored in a markup language document, in a single file dedicated to the program in question, or in multiple coordinated files, e.g., files that store one DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT Applicationor more modules, sub programs, or portions of code. A computer program can be deployed to be executed on one computer or on multiple computers that are located at one site or distributed across multiple sites and interconnected by a data communication network. In this specification, the term “database” is used broadly to refer to any collection of data: the data does not need to be structured in any particular way, or structured at all, and it can be stored on storage devices in one or more locations. Thus, for example, the index database can include multiple collections of data, each of which may be organized and accessed differently. Similarly, in this specification the term “engine” is used broadly to refer to a software-based system, subsystem, or process that is programmed to perform one or more specific functions. Generally, an engine will be implemented as one or more software modules or components, installed on one or more computers in one or more locations. In some cases, one or more computers will be dedicated to a particular engine; in other cases, multiple engines can be installed and running on the same computer or computers. The processes and logic flows described in this specification can be performed by one or more programmable computers executing one or more computer programs to perform functions by operating on input data and generating output. The processes and logic flows can also be performed by special purpose logic circuitry, e.g., an FPGA or an ASIC, or by a combination of special purpose logic circuitry and one or more programmed computers. Computers suitable for the execution of a computer program can be based on general or special purpose microprocessors or both, or any other kind of central processing unit. Generally, a central processing unit will receive instructions and data from a read only memory or a random access memory or both. The essential elements of a computer are a central processing unit for performing or executing instructions and one or more memory devices for storing instructions and data. The central processing unit and the memory can be supplemented by, or incorporated in, special purpose logic circuitry. Generally, a computer will also include, or be operatively coupled to receive data from or transfer data to, or both, one or more mass storage devices for storing data, e.g., magnetic, magneto optical disks, or optical disks. However, a computer need not have such devices. Moreover, a computer can be embedded in another device, e.g., a mobile telephone, a personal digital assistant (PDA), a mobile audio or video player, a game console, a Global Positioning System (GPS) receiver, or a portable storage device, e.g., a universal serial bus (USB) flash drive, to name just a few. Computer readable media suitable for storing computer program instructions and data include all forms of non-volatile memory, media and memory devices, including by way of DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT Applicationexample semiconductor memory devices, e.g., EPROM, EEPROM, and flash memory devices; magnetic disks, e.g., internal hard disks or removable disks; magneto optical disks; and CD ROM and DVD-ROM disks. To provide for interaction with a user, embodiments of the subject matter described in this specification can be implemented on a computer having a display device, e.g., a CRT (cathode ray tube) or LCD (liquid crystal display) monitor, for displaying information to the user and a keyboard and a pointing device, e.g., a mouse or a trackball, by which the user can provide input to the computer. Other kinds of devices can be used to provide for interaction with a user as well; for example, feedback provided to the user can be any form of sensory feedback, e.g., visual feedback, auditory feedback, or tactile feedback; and input from the user can be received in any form, including acoustic, speech, or tactile input. In addition, a computer can interact with a user by sending documents to and receiving documents from a device that is used by the user; for example, by sending web pages to a web browser on a user’s device in response to requests received from the web browser. Also, a computer can interact with a user by sending text messages or other forms of message to a personal device, e.g., a smartphone that is running a messaging application, and receiving responsive messages from the user in return. Data processing apparatus for implementing machine learning models can also include, for example, special-purpose hardware accelerator units for processing common and compute-intensive parts of machine learning training or production, e.g., inference, workloads. Machine learning models can be implemented and deployed using a machine learning framework, .e.g., a TensorFlow framework or a Jax framework. Embodiments of the subject matter described in this specification can be implemented in a computing system that includes a back end component, e.g., as a data server, or that includes a middleware component, e.g., an application server, or that includes a front end component, e.g., a client computer having a graphical user interface, a web browser, or an app through which a user can interact with an implementation of the subject matter described in this specification, or any combination of one or more such back end, middleware, or front end components. The components of the system can be interconnected by any form or medium of digital data communication, e.g., a communication network. Examples of communication networks include a local area network (LAN) and a wide area network (WAN), e.g., the Internet. DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationThe computing system can include clients and servers. A client and server are generally remote from each other and typically interact through a communication network. The relationship of client and server arises by virtue of computer programs running on the respective computers and having a client-server relationship to each other. In some embodiments, a server transmits data, e.g., an HTML page, to a user device, e.g., for purposes of displaying data to and receiving user input from a user interacting with the device, which acts as a client. Data generated at the user device, e.g., a result of the user interaction, can be received at the server from the device. While this specification contains many specific implementation details, these should not be construed as limitations on the scope of any invention or on the scope of what may be claimed, but rather as descriptions of features that may be specific to particular embodiments of particular inventions. Certain features that are described in this specification in the context of separate embodiments can also be implemented in combination in a single embodiment. Conversely, various features that are described in the context of a single embodiment can also be implemented in multiple embodiments separately or in any suitable subcombination. Moreover, although features may be described above as acting in certain combinations and even initially be claimed as such, one or more features from a claimed combination can in some cases be excised from the combination, and the claimed combination may be directed to a subcombination or variation of a subcombination. Similarly, while operations are depicted in the drawings and recited in the claims in a particular order, this should not be understood as requiring that such operations be performed in the particular order shown or in sequential order, or that all illustrated operations be performed, to achieve desirable results. In certain circumstances, multitasking and parallel processing may be advantageous. Moreover, the separation of various system modules and components in the embodiments described above should not be understood as requiring such separation in all embodiments, and it should be understood that the described program components and systems can generally be integrated together in a single software product or packaged into multiple software products. Particular embodiments of the subject matter have been described. Other embodiments are within the scope of the following claims. For example, the actions recited in the claims can be performed in a different order and still achieve desirable results. As one example, the processes depicted in the accompanying figures do not necessarily require the particular order shown, or sequential order, to achieve desirable results. In some cases, multitasking and parallel processing may be advantageous. DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationAspects of the present disclosure may be as set out in the following clauses: Clause 1. A method performed by one or more computers, the method comprising: receiving mesh data specifying a configuration of a plurality of hardware devices that assigns each of the plurality of hardware devices to a respective index along each of one or more axes; obtaining model specification data characterizing a training step for training a neural network, the specification data identifying a plurality of tensors, the plurality of tensors comprising: (i) a plurality of tensors that are provided as input to the training step; and (ii) one or more tensors that are generated as output of the training step; obtaining a schedule for partitioning the training of a neural network across the plurality of hardware devices, wherein the schedule comprises a sequence of partitioning tactics, and wherein each partitioning tactic defines, for each of one or more tensors of the plurality of tensors, a partitioning of a respective dimension of the tensor across one of the one or more axes; and processing the mesh data, the model specification data, and the schedule to generate program data that, when executed by each of the plurality of hardware devices, partitions the training step across the plurality of hardware devices. Clause 2. The method of clause 1, further comprising: providing the program data to a compiler for compilation. Clause 3. The method of clause 2, further comprising: causing the plurality of hardware devices to execute compiled program data generated by the compiler in order to perform one or more training steps for training the neural network. Clause 4. The method of any preceding clause, wherein when performing the training step, the plurality of hardware devices operate under a single program multiple data (SPMD) model and each execute a same program, and wherein the program data specifies the same program to be executed by each of the plurality of devices to perform the training step. Clause 5. The method of clause 4, wherein the program data comprises device-local DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationSPMD code for execution by each of the plurality of hardware devices. Clause 6. The method of any preceding clause, wherein the partitioning tactics in the sequence comprise: one or more manual partitioning tactics that identify, for each of one or more tensors of the plurality of tensors, the respective dimension of the tensor to be partitioned across the axis. Clause 7. The method of any preceding clause, wherein the partitioning tactics in the sequence comprise: one or more automatic partitioning tactics that identify a respective axis but do not identify the respective dimension of the tensor to be partitioned across the respective axis. Clause 8. The method of clause 7, wherein processing the mesh data, the model specification data, and the schedule to generate program data that, when executed by each of the plurality of hardware devices, partitions the training step across the plurality of hardware devices comprises, for each automatic partitioning tactic, performing, using a runtime cost model, a search to identify a particular dimension of a particular tensor to be partitioned across the axis identified by the automatic partitioning tactic. Clause 9. The method of clause 7, when dependent on clause 6, wherein the partitioning tactics comprise one or more manual partitioning tactics and one or more automatic partitioning tactics. Clause 10. The method of any preceding clause, wherein processing the mesh data, the model specification data, and the schedule comprises: mapping each partitioning tactic in the schedule to a sequence of compiler actions. Clause 11. The method of clause 10, wherein processing the mesh data, the model specification data, and the schedule comprises: propagating the compiler actions through the program data based on linear algebra homomorphisms defined by the partitioning tactics in the schedule. DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationClause 12. The method of any preceding clause, wherein processing the mesh data, the model specification data, and the schedule comprises: incrementally resolving conflicts caused by partitioning tactics according to an order of the partitioning tactics within the sequence. Clause 13. A system comprising one or more computers and one or more storage devices storing instructions that when executed by the one or more computers cause the one more computers to perform the operations of the respective method of any one of clauses 1-12. Clause 14. One or more computer storage media storing instructions that when executed by one or more computers cause the one more computers to perform the operations of the respective method of any one of clauses 1-12.
Claims
DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT ApplicationCLAIMS 1. A method performed by one or more computers, the method comprising: receiving mesh data specifying a configuration of a plurality of hardware devices that assigns each of the plurality of hardware devices to a respective index along each of one or more axes; obtaining model specification data characterizing a training step for training a neural network, the specification data identifying a plurality of tensors, the plurality of tensors comprising: (i) a plurality of tensors that are provided as input to the training step; and (ii) one or more tensors that are generated as output of the training step; obtaining a schedule for partitioning the training of a neural network across the plurality of hardware devices, wherein the schedule comprises a sequence of partitioning tactics, and wherein each partitioning tactic defines, for each of one or more tensors of the plurality of tensors, a partitioning of a respective dimension of the tensor across one of the one or more axes; and processing the mesh data, the model specification data, and the schedule to generate program data that, when executed by each of the plurality of hardware devices, partitions the training step across the plurality of hardware devices.
2. The method of claim 1, further comprising: providing the program data to a compiler for compilation.
3. The method of claim 2, further comprising: causing the plurality of hardware devices to execute compiled program data generated by the compiler in order to perform one or more training steps for training the neural network.
4. The method of any preceding claim, wherein when performing the training step, the plurality of hardware devices operate under a single program multiple data (SPMD) model and each execute a same program, and wherein the program data specifies the same program to be executed by each of the plurality of devices to perform the training step.DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT Application5. The method of claim 4, wherein the program data comprises device-local SPMD code for execution by each of the plurality of hardware devices.
6. The method of any preceding claim, wherein the partitioning tactics in the sequence comprise: one or more manual partitioning tactics that identify, for each of one or more tensors of the plurality of tensors, the respective dimension of the tensor to be partitioned across the axis.
7. The method of any preceding claim, wherein the partitioning tactics in the sequence comprise: one or more automatic partitioning tactics that identify a respective axis but do not identify the respective dimension of the tensor to be partitioned across the respective axis.
8. The method of claim 7, wherein processing the mesh data, the model specification data, and the schedule to generate program data that, when executed by each of the plurality of hardware devices, partitions the training step across the plurality of hardware devices comprises, for each automatic partitioning tactic, performing, using a runtime cost model, a search to identify a particular dimension of a particular tensor to be partitioned across the axis identified by the automatic partitioning tactic.
9. The method of claim 7, when dependent on claim 6, wherein the partitioning tactics comprise one or more manual partitioning tactics and one or more automatic partitioning tactics.
10. The method of any preceding claim, wherein processing the mesh data, the model specification data, and the schedule comprises: mapping each partitioning tactic in the schedule to a sequence of compiler actions.
11. The method of claim 10, wherein processing the mesh data, the model specification data, and the schedule comprises: propagating the compiler actions through the program data based on linear algebra homomorphisms defined by the partitioning tactics in the schedule.DeepMind Technologies LimitedF&R Ref.: 45288-0415WO1 PCT Application12. The method of any preceding claim, wherein processing the mesh data, the model specification data, and the schedule comprises: incrementally resolving conflicts caused by partitioning tactics according to an order of the partitioning tactics within the sequence.
13. A system comprising one or more computers and one or more storage devices storing instructions that when executed by the one or more computers cause the one more computers to perform the operations of the respective method of any one of claims 1-12.
14. One or more computer storage media storing instructions that when executed by one or more computers cause the one more computers to perform the operations of the respective method of any one of claims 1-12.