Concepts / Multi-Input Models

Multi-Input Models

A Keras functional model can share one representation and branch into several named prediction heads.

  • Programming

One Graph, Several Predictions

A multi-input or multi-output model is easier to understand when you stop imagining it as one straight stack of layers. A Keras functional model can receive several streams of data, process each stream through its own branch, and merge the resulting representations. It can also share one representation and branch into several named prediction heads. The graph therefore has places where information is separate and places where it is shared.

feedsfeedsfeedsSharedrepresentationcommon featuresAge headage predictionIncome headincome-group predictionGender headgender prediction
How can one shared representation split into multiple named predictions, and which output does each prediction head produce?

Tracing Two Input Branches

A single-input model receives one stream of data. A multi-input model receives several streams and gives each stream its own path before combining the resulting representations. For example, a question-answering system can receive a question through one input and a text snippet through another. The question and text travel independently until a merge operation brings their branch representations together.

travels throughtravels throughrepresentationrepresentationfeedstextinput nodeText encodertext representationMergecombined representationAnsweroutputquestioninput nodeQuestion encoderquestion representation
What happens to separate inputs as they pass through their branch encoders, merge, and reach the output layer?

Following the Question and Text

Trace the intended path for a question-answering model with named inputs text and question.

Input assignment: The text data belongs to the text input node, and the question data belongs to the question input node.

Branch processing: Each input travels through its own branch encoder, producing a representation for that source.

Merge: The two branch representations meet through a Keras merge operation.

Prediction: The merged representation travels to the answer output.

The intended path is text to text branch, question to question branch, both branches to merge, and merge to answer.

Building the Branch Structure

The main construction stages of a two-input functional model are input layers, independent branch encoders, a merge operation, and an output layer. Input names such as text and question identify the input nodes. Those names are useful later because dictionary-based training uses them to associate each array with the correct branch. The model construction is explicit: it receives both input tensors and produces the answer tensor.

The branches meet through a Keras merge operation such as concatenate or add. The merge line is not a minor implementation detail. It determines how the branch representations are brought together, so it is one of the first places to inspect when the model behaves unexpectedly.

branch outputbranch outputbranch outputbranch outputproducesproducesBranch ArepresentationConcatenatecombined representationJoined featuresconcatenated resultBranch BrepresentationAddcombined representationSummed featuresadded result
What changes in the merged representation when two branch outputs are concatenated instead of added element by element?
Merge strategyWhat it doesInspection question
ConcatenateBrings the branch representations together as a combined representationDid the merge preserve the separate branch information in the intended arrangement?
AddCombines the branch representations through an addition operationAre the branch outputs suitable for the intended addition?

Routing Training Data

Training data must preserve the model's input structure. With list-based training, input arrays are supplied in the same order as the model's inputs. This form is concise, but the order must remain correct. With dictionary-based training, arrays are supplied under the names of the input nodes. This makes the pairing explicit and requires named inputs.

input orderinput orderinput nameinput nameList item 1text arraytexttext arraytext inputnamed branchList item 2question arrayquestionquestion arrayquestion inputnamed branch
How does each input array travel to the correct branch when training with a list of arrays versus a dictionary of named arrays?

Choosing a Training Form

A model has named inputs text and question. Organize the two input arrays for training.

List form: Place the text array first and the question array second when that is the model's input order.

Dictionary form: Associate the text array with the key text and the question array with the key question.

Verification: Check that the selected form agrees with the model's input order or input names before training.

Both forms can preserve the intended mapping, but the list depends on order while the dictionary makes the names explicit.

Matching Heads to Losses

A multi-output model makes several predictions from the same input or shared internal representation. Each output can represent a different task, so each head needs a loss suited to what it predicts. The source example uses an age head for scalar regression, an income head for classification over income groups, and a gender head for binary classification. These tasks therefore use different kinds of losses rather than one loss reused blindly for every output.

Prediction headTask described in the sourceOutput behaviorSuitable loss examples
AgeScalar regressionOne unitMean squared error
IncomeClassification over income groupsOne unit per income group with softmaxCategorical crossentropy
GenderBinary classificationOne unit with sigmoidBinary crossentropy

The output task determines the kind of loss that should be associated with the head.

During training, one target array is paired with each output. Keras accepts those targets as a list ordered like the model outputs or as a dictionary keyed by output names. The dictionary form makes the pairing explicit; the list form is shorter but depends on keeping the output order correct.

paired withpaired withpaired withevaluated byevaluated byevaluated byAgeprediction headAge targettarget arrayMSEage lossIncomeprediction headIncome targettarget arrayCategoricalcrossentropyincome lossGenderprediction headGender targettarget arrayBinary crossentropygender loss
How does each prediction head connect to its corresponding target and loss?

Forming the Global Loss

The separate output losses do not remain isolated during optimization. Keras combines the individual losses into one global loss used to optimize the model. Because the output heads share a representation, the combined result determines how the shared part of the network is trained.

A plain combination does not guarantee that every task influences learning equally. If one individual loss is numerically much larger than the others, optimization can be driven mainly by that task. Loss weights change each loss's contribution before the global loss is formed, helping balance outputs whose losses have different numerical scales.

contributioncontributioncontributionformsoptimizesAge lossindividual lossLoss weightsoptional balancingGlobal lossoptimization signalShared representationupdated modelIncome lossindividual lossGender lossindividual loss
How do individual losses from several outputs become one global loss used to update the model?

Balancing Three Tasks

A shared representation feeds age, income, and gender heads. The age loss is numerically much larger than the other two losses. What should you inspect?

Identify the imbalance: Compare the numerical scales of the individual losses rather than assuming that their unweighted contributions are equally influential.

Inspect the global combination: Keras combines the individual losses into one global loss, so a larger individual loss can have a larger effect on optimization.

Consider loss weights: Loss weights change each loss's contribution before the global loss is formed and can help prevent the age task from dominating the shared representation.

The likely inspection point is the loss-weight configuration and the relative numerical scales of the individual losses.

Debugging the Data Path

A multi-input model can fail at several distinct boundaries. Treat it as a trace: data enters an input node, passes through its branch, reaches the merge operation, and contributes to the output. For a multi-output model, trace each prediction head to its loss, optional weight, and target data. The most useful debugging question is where the actual path first differs from the intended graph.

mapping agreesbranch agreesmerge agreesall paths agreeData mappingorder or namesInput branchintended data pathMerge stagebranch combinationOutput mappingtarget and loss pathIntended graphtrace continues
How can a failure be localized by checking whether it occurs in an input branch, at the merge operation, or in the mapping between data and named inputs or outputs?
  • Supplying list-based inputs in an order different from the model's input order.

    List-based training pairs arrays by position, so a changed order sends data to the wrong branch.

    Fix: Check the model's input order before supplying a list, or use named inputs with a dictionary.

  • Using dictionary-based inputs without matching the model's input names.

    Dictionary-based training relies on input names to make the pairing explicit.

    Fix: Compare every dictionary key with the model's input names.

  • Treating every output as if it represented the same task.

    The output heads represent different prediction tasks and need losses suited to those tasks.

    Fix: Trace each head to its task-appropriate loss.

  • Assuming a plain combination of losses gives every task equal influence.

    Different numerical scales can make one loss contribute more strongly to the global optimization signal.

    Fix: Inspect the loss scales and consider loss weights.

  • Inspecting only the final prediction when a multi-input model behaves unexpectedly.

    The model is a graph with several boundaries, not one undifferentiated object.

    Fix: Trace data mapping, each input branch, the merge operation, and the output path separately.

Practice the Trace

MEDIUM

A functional model has two named inputs, text and question, and one answer output. Explain how you would verify the data path before training. Then describe how the training data would differ between list-based and dictionary-based input forms. Finally, imagine that the model has age, income, and gender outputs instead. State what you would inspect for each output besides the prediction itself.

Hints
  • Begin with the mapping between arrays and input names or input order.
  • Follow each input through its independent branch and inspect the merge operation.
  • For several outputs, trace each head to its target, loss, and optional loss weight.
  1. A multi-input Keras model gives each data source an independent branch before merging the branch representations. Named input nodes support dictionary-based training, while list-based training depends on preserving input order. A merge operation such as concatenate or add determines how branch outputs meet. A multi-output model can reuse a shared representation for several named prediction heads. Each head needs a task-appropriate loss and target data, and Keras combines the individual losses into one global optimization loss. Loss weights help when one loss has a much larger numerical scale. When debugging, trace the actual data path from mapping to branch, merge, output, target, loss, and optional weight.

Key Takeaways

  • Separate inputs travel through independent branches and meet at a Keras merge operation.
  • Concatenate and add are different merge strategies, so the merge stage deserves explicit inspection.
  • One shared representation can feed several named prediction heads for different tasks.
  • Each output needs matching target data and a loss suited to its task; loss weights can balance different loss scales.
  • List-based training depends on order, dictionary-based training depends on names, and debugging should trace the complete data path.