Introduction and Implementation of DBMTL for Multi-task Learning Model
マルチタスク学習の背景
現在業界で使用されているレコメンデーションアルゴリズムは、単一ターゲット(CTR)タスクに限らず、コメント、ブックマーク、カート追加、購入、視聴時間などの後続コンバージョンリンクにも注目する必要があります。
一般的な多目的最適化モデルは、各最適化目標に対して独立したモデルネットワークから出発し、これらのネットワークが下位層でパラメータを共有することで、各目標に関連するモデルの適切な独立性和相関性を実現します。このタイプのモデルフレームワークは、上記の図の構造で概括できます。下位層のパラメータ共有方法に関わらず、これらのネットワークは最後の数層で独立した分岐を持ち、各ターゲットの最終値を予測します。このようなネットワークの確率モデルは次の数式で記述できます。
ここで、l と m はターゲット、x はサンプル特徴、H はモデルです。ここでは各ターゲットが独立であるという仮定を置いています。
DBMTL の概要
DBMTL(Deep Bayesian Multi-Target Learning、深層ベイズ多目標学習)の出発点の一つは、上記の問題を解決することです。単純なベイズの公式を適用すると、確率モデルは次のように記述できます。
下の図に示すように、DBMTL と従来の MTL 構造(各ターゲットを独立とみなす)の主な違いは、ターゲットノード間にベイジアンネットワークを構築し、ターゲット間の因果関係を明示的にモデル化する点にあります。実際のビジネスでは、ユーザーの行動には明確な逐次的依存関係があることが多いためです。たとえば、情報フィードシナリオでは、ユーザーはまず記事の詳細ページをクリックしてから、閲覧、コメント、転送、お気に入りなどの後続操作を行います。DBMTL はこれらの関係をモデル構造に組み込んでおり、より優れた学習結果が期待できます。
次の図は、DBMTL モデルの具体的な実装です。ネットワークは入力層、共有埋め込み層、共有層、識別層、ベイズ層で構成されます。
・共有埋め込み層は、各ターゲットのトレーニングで共有されるルックアップテーブルです。
・共有層と分割層は、ターゲットの共通表現と固有表現をそれぞれモデル化する汎用の多層パーセプトロン(MLP)です。
・ベイズ層は DBMTL の最も重要な部分です。次の確率モデルを実装します。
対応する対数尤度損失関数は次の通りです。
実用上、異なる目標に対する重み調整は依然として大きな実用効果があります。ターゲットに異なる重みを割り当てる場合、損失関数を次のように再定式化することに相当します。
ネットワークのベイズ層では、関数 f1、f2、f3 はターゲット間の暗黙の因果関係を学習するために全結合 MLP として実装されます。関数の入力変数の埋め込みを連結したものを入力とし、関数の出力変数を表す埋め込みを出力します。各ターゲットの埋め込みは最終的に MLP の一層を通じて、最終ターゲットの確率を出力します。
コード
EasyRec レコメンデーションアルゴリズムフレームワークに基づき、DBMTL アルゴリズムを実装しました。具体的な実装は GitHub の EasyRec-DBMTL で確認できます。
EasyRec の概要:EasyRec は Alibaba Cloud コンピューティングプラットフォームの機械学習 PAI チームがオープンソース化した大規模分散レコメンデーションアルゴリズムフレームワークです。優れた効果を発揮する特徴量エンジニアリング手法、統合トレーニング、評価、デプロイメントを備え、Alibaba Cloud プロダクトとシームレスに連携し、短期間で最先端のレコメンデーションシステムを構築できます。Alibaba Cloud の先導プロダクトとして、数百の企業顧客に安定的にサービスを提供してきました。
モデルフィードフォワードネットワーク
def build_predict_graph(self):
"""Forward function.
Returns:
self._prediction_dict: Prediction result of two tasks.
"""
# Here we start from the tensor (self._features) after sharing the embedding layer, omitting its generation logic
# shared layer
if self._model_config.HasField('bottom_dnn'):
bottom_dnn = dnn.DNN(
self._model_config.bottom_dnn,
self._l2_reg,
name='bottom_dnn',
is_training=self._is_training)
bottom_fea = bottom_dnn(self._features)
else:
bottom_fea = self._features
# MMOE block
if self._model_config.HasField('expert_dnn'):
mmoe_layer = mmoe.MMOE(
self._model_config.expert_dnn,
l2_reg=self._l2_reg,
num_task=self._task_num,
num_expert=self._model_config.num_expert)
task_input_list = mmoe_layer(bottom_fea)
else:
task_input_list = [bottom_fea] * self._task_num
tower_features = {}
# specific layer
for i, task_tower_cfg in enumerate(self._model_config.task_towers):
tower_name = task_tower_cfg.tower_name
if task_tower_cfg. HasField('dnn'):
tower_dnn = dnn.DNN(
task_tower_cfg.dnn,
self._l2_reg,
name=tower_name + '/dnn',
is_training=self._is_training)
tower_fea = tower_dnn(task_input_list[i])
tower_features[tower_name] = tower_fea
else:
tower_features[tower_name] = task_input_list[i]
tower_outputs = {}
relation_features = {}
#bayesian network
for task_tower_cfg in self._model_config.task_towers:
tower_name = task_tower_cfg.tower_name
relation_dnn = dnn.DNN(
task_tower_cfg.relation_dnn,
self._l2_reg,
name=tower_name + '/relation_dnn',
is_training=self._is_training)
tower_inputs = [tower_features[tower_name]]
for relation_tower_name in task_tower_cfg.relation_tower_names:
tower_inputs.append(relation_features[relation_tower_name])
relation_input = tf.concat(
tower_inputs, axis=-1, name=tower_name + '/relation_input')
relation_fea = relation_dnn(relation_input)
relation_features[tower_name] = relation_features
output_logits = tf.layers.dense(
relation_fea,
task_tower_cfg.num_class,
kernel_regularizer=self._l2_reg,
name=tower_name + '/output')
tower_outputs[tower_name] = output_logits
self._add_to_prediction_dict(tower_outputs)
損失計算
def build(loss_type, label, pred, loss_weight=1.0, num_class=1, **kwargs):
if loss_type == LossType. CLASSIFICATION:
if num_class == 1:
return tf.losses.sigmoid_cross_entropy(
label, logits=pred, weights=loss_weight, **kwargs)
else:
return tf.losses.sparse_softmax_cross_entropy(
labels=label, logits=pred, weights=loss_weight, **kwargs)
elif loss_type == LossType.CROSS_ENTROPY_LOSS:
return tf.losses.log_loss(label, pred, weights=loss_weight, **kwargs)
elif loss_type in [LossType.L2_LOSS, LossType.SIGMOID_L2_LOSS]:
logging.info('%s is used' % LossType.Name(loss_type))
return tf.losses.mean_squared_error(
labels=label, predictions=pred, weights=loss_weight, **kwargs)
elif loss_type == LossType. PAIR_WISE_LOSS:
return pairwise_loss(pred, label)
else:
raise ValueError('unsupported loss type: %s' % LossType.Name(loss_type))
def _build_loss_impl(self,
loss_type,
label_name,
loss_weight=1.0,
num_class=1,
suffix=''):
loss_dict = {}
if loss_type == LossType. CLASSIFICATION:
loss_name = 'cross_entropy_loss' + suffix
pred = self._prediction_dict['logits' + suffix]
elif loss_type in [LossType.L2_LOSS, LossType.SIGMOID_L2_LOSS]:
loss_name = 'l2_loss' + suffix
pred = self._prediction_dict['y' + suffix]
else:
raise ValueError('invalid loss type: %s' % LossType.Name(loss_type))
loss_dict[loss_name] = build(loss_type,
self._labels[label_name],
pred,
loss_weight, num_class)
return loss_dict
def build_loss_graph(self):
"""Build loss graph for multi task model."""
for task_tower_cfg in self._task_towers:
tower_name = task_tower_cfg.tower_name
loss_weight = task_tower_cfg.weight * self._sample_weight
if hasattr(task_tower_cfg, 'task_space_indicator_label') and
task_tower_cfg. HasField('task_space_indicator_label'):
in_task_space = tf.to_float(
self._labels[task_tower_cfg.task_space_indicator_label] > 0)
loss_weight = loss_weight * (
task_tower_cfg.in_task_space_weight * in_task_space +
task_tower_cfg.out_task_space_weight * (1 - in_task_space))
# The EasyRec framework will automatically add the loss in self._loss_dict.
self._loss_dict.update(
self._build_loss_impl(
task_tower_cfg.loss_type,
label_name=self._label_name_dict[tower_name],
loss_weight=loss_weight,
num_class=task_tower_cfg.num_class,
suffix='_%s' % tower_name))
return self._loss_dict
アプリケーション
DBMTL は優れたアルゴリズム効果により、PAI プラットフォームで広く利用されています。
ライブストリーミングレコメンデーションを例にとると、このシナリオには is_click、is_view、view_costtime、is_on_mic、on_mic_duration の複数の目標があり、そのうち is_click、is_view、is_on_mic は二項分類タスク、view_costtime と on_mic_duration は持続時間を予測する回帰タスクです。ユーザー行動の依存関係は次の通りです。
・is_click => is_view
・is_click+is_view=> view_costtime
・is_click => is_on_mic
・is_click+is_on_mic => on_mic_duration
したがって、設定は次のようになります。
dbmtl {
bottom_dnn {
hidden_units: [512, 256]
}
task_towers {
tower_name: "is_click"
label_name: "is_click"
loss_type: CLASSIFICATION
metrics_set: {
auc {}
}
dnn {
hidden_units: [128, 96, 64]
}
relation_dnn {
hidden_units: [32]
}
weight: 1.0
}
task_towers {
tower_name: "is_view"
label_name: "is_view"
loss_type: CLASSIFICATION
metrics_set: {
auc {}
}
dnn {
hidden_units: [128, 96, 64]
}
relation_tower_names: ["is_click"]
relation_dnn {
hidden_units: [32]
}
weight: 1.0
}
task_towers {
tower_name: "view_costtime"
label_name: "view_costtime"
loss_type: L2_LOSS
metrics_set: {
mean_squared_error {}
}
dnn {
hidden_units: [128, 96, 64]
}
relation_tower_names: ["is_click", "is_view"]
relation_dnn {
hidden_units: [32]
}
weight: 1.0
}
task_towers {
tower_name: "is_on_mic"
label_name: "is_on_mic"
loss_type: CLASSIFICATION
metrics_set: {
auc {}
}
dnn {
hidden_units: [128, 96, 64]
}
relation_tower_names: ["is_click"]
relation_dnn {
hidden_units: [32]
}
weight: 1.0
}
task_towers {
tower_name: "on_mic_duration"
label_name: "on_mic_duration"
loss_type: L2_LOSS
metrics_set: {
mean_squared_error {}
}
dnn {
hidden_units: [128, 96, 64]
}
relation_tower_names: ["is_click", "is_on_mic"]
relation_dnn {
hidden_units: [32]
}
weight: 1.0
}
l2_regularization: 1e-6
}
embedding_regularization: 5e-6
}
特筆すべきは、DBMTL モデルの導入後、GBDT+FM(傍観単一ターゲット)と比較して、オンライン視聴率が 18% 向上し、マイク接続率が 14% 向上したことです。
現在業界で使用されているレコメンデーションアルゴリズムは、単一ターゲット(CTR)タスクに限らず、コメント、ブックマーク、カート追加、購入、視聴時間などの後続コンバージョンリンクにも注目する必要があります。
一般的な多目的最適化モデルは、各最適化目標に対して独立したモデルネットワークから出発し、これらのネットワークが下位層でパラメータを共有することで、各目標に関連するモデルの適切な独立性和相関性を実現します。このタイプのモデルフレームワークは、上記の図の構造で概括できます。下位層のパラメータ共有方法に関わらず、これらのネットワークは最後の数層で独立した分岐を持ち、各ターゲットの最終値を予測します。このようなネットワークの確率モデルは次の数式で記述できます。
ここで、l と m はターゲット、x はサンプル特徴、H はモデルです。ここでは各ターゲットが独立であるという仮定を置いています。
DBMTL の概要
DBMTL(Deep Bayesian Multi-Target Learning、深層ベイズ多目標学習)の出発点の一つは、上記の問題を解決することです。単純なベイズの公式を適用すると、確率モデルは次のように記述できます。
下の図に示すように、DBMTL と従来の MTL 構造(各ターゲットを独立とみなす)の主な違いは、ターゲットノード間にベイジアンネットワークを構築し、ターゲット間の因果関係を明示的にモデル化する点にあります。実際のビジネスでは、ユーザーの行動には明確な逐次的依存関係があることが多いためです。たとえば、情報フィードシナリオでは、ユーザーはまず記事の詳細ページをクリックしてから、閲覧、コメント、転送、お気に入りなどの後続操作を行います。DBMTL はこれらの関係をモデル構造に組み込んでおり、より優れた学習結果が期待できます。
次の図は、DBMTL モデルの具体的な実装です。ネットワークは入力層、共有埋め込み層、共有層、識別層、ベイズ層で構成されます。
・共有埋め込み層は、各ターゲットのトレーニングで共有されるルックアップテーブルです。
・共有層と分割層は、ターゲットの共通表現と固有表現をそれぞれモデル化する汎用の多層パーセプトロン(MLP)です。
・ベイズ層は DBMTL の最も重要な部分です。次の確率モデルを実装します。
対応する対数尤度損失関数は次の通りです。
実用上、異なる目標に対する重み調整は依然として大きな実用効果があります。ターゲットに異なる重みを割り当てる場合、損失関数を次のように再定式化することに相当します。
ネットワークのベイズ層では、関数 f1、f2、f3 はターゲット間の暗黙の因果関係を学習するために全結合 MLP として実装されます。関数の入力変数の埋め込みを連結したものを入力とし、関数の出力変数を表す埋め込みを出力します。各ターゲットの埋め込みは最終的に MLP の一層を通じて、最終ターゲットの確率を出力します。
コード
EasyRec レコメンデーションアルゴリズムフレームワークに基づき、DBMTL アルゴリズムを実装しました。具体的な実装は GitHub の EasyRec-DBMTL で確認できます。
EasyRec の概要:EasyRec は Alibaba Cloud コンピューティングプラットフォームの機械学習 PAI チームがオープンソース化した大規模分散レコメンデーションアルゴリズムフレームワークです。優れた効果を発揮する特徴量エンジニアリング手法、統合トレーニング、評価、デプロイメントを備え、Alibaba Cloud プロダクトとシームレスに連携し、短期間で最先端のレコメンデーションシステムを構築できます。Alibaba Cloud の先導プロダクトとして、数百の企業顧客に安定的にサービスを提供してきました。
モデルフィードフォワードネットワーク
def build_predict_graph(self):
"""Forward function.
Returns:
self._prediction_dict: Prediction result of two tasks.
"""
# Here we start from the tensor (self._features) after sharing the embedding layer, omitting its generation logic
# shared layer
if self._model_config.HasField('bottom_dnn'):
bottom_dnn = dnn.DNN(
self._model_config.bottom_dnn,
self._l2_reg,
name='bottom_dnn',
is_training=self._is_training)
bottom_fea = bottom_dnn(self._features)
else:
bottom_fea = self._features
# MMOE block
if self._model_config.HasField('expert_dnn'):
mmoe_layer = mmoe.MMOE(
self._model_config.expert_dnn,
l2_reg=self._l2_reg,
num_task=self._task_num,
num_expert=self._model_config.num_expert)
task_input_list = mmoe_layer(bottom_fea)
else:
task_input_list = [bottom_fea] * self._task_num
tower_features = {}
# specific layer
for i, task_tower_cfg in enumerate(self._model_config.task_towers):
tower_name = task_tower_cfg.tower_name
if task_tower_cfg. HasField('dnn'):
tower_dnn = dnn.DNN(
task_tower_cfg.dnn,
self._l2_reg,
name=tower_name + '/dnn',
is_training=self._is_training)
tower_fea = tower_dnn(task_input_list[i])
tower_features[tower_name] = tower_fea
else:
tower_features[tower_name] = task_input_list[i]
tower_outputs = {}
relation_features = {}
#bayesian network
for task_tower_cfg in self._model_config.task_towers:
tower_name = task_tower_cfg.tower_name
relation_dnn = dnn.DNN(
task_tower_cfg.relation_dnn,
self._l2_reg,
name=tower_name + '/relation_dnn',
is_training=self._is_training)
tower_inputs = [tower_features[tower_name]]
for relation_tower_name in task_tower_cfg.relation_tower_names:
tower_inputs.append(relation_features[relation_tower_name])
relation_input = tf.concat(
tower_inputs, axis=-1, name=tower_name + '/relation_input')
relation_fea = relation_dnn(relation_input)
relation_features[tower_name] = relation_features
output_logits = tf.layers.dense(
relation_fea,
task_tower_cfg.num_class,
kernel_regularizer=self._l2_reg,
name=tower_name + '/output')
tower_outputs[tower_name] = output_logits
self._add_to_prediction_dict(tower_outputs)
損失計算
def build(loss_type, label, pred, loss_weight=1.0, num_class=1, **kwargs):
if loss_type == LossType. CLASSIFICATION:
if num_class == 1:
return tf.losses.sigmoid_cross_entropy(
label, logits=pred, weights=loss_weight, **kwargs)
else:
return tf.losses.sparse_softmax_cross_entropy(
labels=label, logits=pred, weights=loss_weight, **kwargs)
elif loss_type == LossType.CROSS_ENTROPY_LOSS:
return tf.losses.log_loss(label, pred, weights=loss_weight, **kwargs)
elif loss_type in [LossType.L2_LOSS, LossType.SIGMOID_L2_LOSS]:
logging.info('%s is used' % LossType.Name(loss_type))
return tf.losses.mean_squared_error(
labels=label, predictions=pred, weights=loss_weight, **kwargs)
elif loss_type == LossType. PAIR_WISE_LOSS:
return pairwise_loss(pred, label)
else:
raise ValueError('unsupported loss type: %s' % LossType.Name(loss_type))
def _build_loss_impl(self,
loss_type,
label_name,
loss_weight=1.0,
num_class=1,
suffix=''):
loss_dict = {}
if loss_type == LossType. CLASSIFICATION:
loss_name = 'cross_entropy_loss' + suffix
pred = self._prediction_dict['logits' + suffix]
elif loss_type in [LossType.L2_LOSS, LossType.SIGMOID_L2_LOSS]:
loss_name = 'l2_loss' + suffix
pred = self._prediction_dict['y' + suffix]
else:
raise ValueError('invalid loss type: %s' % LossType.Name(loss_type))
loss_dict[loss_name] = build(loss_type,
self._labels[label_name],
pred,
loss_weight, num_class)
return loss_dict
def build_loss_graph(self):
"""Build loss graph for multi task model."""
for task_tower_cfg in self._task_towers:
tower_name = task_tower_cfg.tower_name
loss_weight = task_tower_cfg.weight * self._sample_weight
if hasattr(task_tower_cfg, 'task_space_indicator_label') and
task_tower_cfg. HasField('task_space_indicator_label'):
in_task_space = tf.to_float(
self._labels[task_tower_cfg.task_space_indicator_label] > 0)
loss_weight = loss_weight * (
task_tower_cfg.in_task_space_weight * in_task_space +
task_tower_cfg.out_task_space_weight * (1 - in_task_space))
# The EasyRec framework will automatically add the loss in self._loss_dict.
self._loss_dict.update(
self._build_loss_impl(
task_tower_cfg.loss_type,
label_name=self._label_name_dict[tower_name],
loss_weight=loss_weight,
num_class=task_tower_cfg.num_class,
suffix='_%s' % tower_name))
return self._loss_dict
アプリケーション
DBMTL は優れたアルゴリズム効果により、PAI プラットフォームで広く利用されています。
ライブストリーミングレコメンデーションを例にとると、このシナリオには is_click、is_view、view_costtime、is_on_mic、on_mic_duration の複数の目標があり、そのうち is_click、is_view、is_on_mic は二項分類タスク、view_costtime と on_mic_duration は持続時間を予測する回帰タスクです。ユーザー行動の依存関係は次の通りです。
・is_click => is_view
・is_click+is_view=> view_costtime
・is_click => is_on_mic
・is_click+is_on_mic => on_mic_duration
したがって、設定は次のようになります。
dbmtl {
bottom_dnn {
hidden_units: [512, 256]
}
task_towers {
tower_name: "is_click"
label_name: "is_click"
loss_type: CLASSIFICATION
metrics_set: {
auc {}
}
dnn {
hidden_units: [128, 96, 64]
}
relation_dnn {
hidden_units: [32]
}
weight: 1.0
}
task_towers {
tower_name: "is_view"
label_name: "is_view"
loss_type: CLASSIFICATION
metrics_set: {
auc {}
}
dnn {
hidden_units: [128, 96, 64]
}
relation_tower_names: ["is_click"]
relation_dnn {
hidden_units: [32]
}
weight: 1.0
}
task_towers {
tower_name: "view_costtime"
label_name: "view_costtime"
loss_type: L2_LOSS
metrics_set: {
mean_squared_error {}
}
dnn {
hidden_units: [128, 96, 64]
}
relation_tower_names: ["is_click", "is_view"]
relation_dnn {
hidden_units: [32]
}
weight: 1.0
}
task_towers {
tower_name: "is_on_mic"
label_name: "is_on_mic"
loss_type: CLASSIFICATION
metrics_set: {
auc {}
}
dnn {
hidden_units: [128, 96, 64]
}
relation_tower_names: ["is_click"]
relation_dnn {
hidden_units: [32]
}
weight: 1.0
}
task_towers {
tower_name: "on_mic_duration"
label_name: "on_mic_duration"
loss_type: L2_LOSS
metrics_set: {
mean_squared_error {}
}
dnn {
hidden_units: [128, 96, 64]
}
relation_tower_names: ["is_click", "is_on_mic"]
relation_dnn {
hidden_units: [32]
}
weight: 1.0
}
l2_regularization: 1e-6
}
embedding_regularization: 5e-6
}
特筆すべきは、DBMTL モデルの導入後、GBDT+FM(傍観単一ターゲット)と比較して、オンライン視聴率が 18% 向上し、マイク接続率が 14% 向上したことです。
Related Articles
-
A detailed explanation of Hadoop core architecture HDFS
Knowledge Base Team
-
What Does IOT Mean
Knowledge Base Team
-
6 Optional Technologies for Data Storage
Knowledge Base Team
-
What Is Blockchain Technology
Knowledge Base Team
Explore More Special Offers
-
Short Message Service(SMS) & Mail Service
50,000 email package starts as low as USD 1.99, 120 short messages start at only USD 1.00
