Add speed feature for jnt_comp

Add two speed features, turned off by default.
(1). flag to indicate full tx type/size search or use model_rd.
(2). flag to indiate whether to skip newmv search

In rd loop, the codec will call motion_mode_rd two times to make
desicion whether to select jnt_comp mode.
Calling motion_mode_rd two times is expensive. Eventually we would
like to substitue full search with limited range search.

Change-Id: I5ec48e795382d44623620119dcebd0ade9322154
diff --git a/av1/encoder/rdopt.c b/av1/encoder/rdopt.c
index 223ed8f..2506710 100644
--- a/av1/encoder/rdopt.c
+++ b/av1/encoder/rdopt.c
@@ -1471,7 +1471,8 @@
                             MACROBLOCK *x, MACROBLOCKD *xd, int plane_from,
                             int plane_to, int *out_rate_sum,
                             int64_t *out_dist_sum, int *skip_txfm_sb,
-                            int64_t *skip_sse_sb) {
+                            int64_t *skip_sse_sb, int *plane_rate,
+                            int64_t *plane_sse, int64_t *plane_dist) {
   // Note our transform coeffs are 8 times an orthogonal transform.
   // Hence quantizer step is also 8 times. To get effective quantizer
   // we need to divide by 8 before sending to modeling function.
@@ -1507,6 +1508,9 @@
 
     rate_sum += rate;
     dist_sum += dist;
+    if (plane_rate) plane_rate[plane] = rate;
+    if (plane_sse) plane_sse[plane] = sse;
+    if (plane_dist) plane_dist[plane] = dist;
   }
 
   *skip_txfm_sb = total_sse == 0;
@@ -3000,7 +3004,8 @@
   }
   // RD estimation.
   model_rd_for_sb(cpi, bsize, x, xd, 0, 0, &this_rd_stats.rate,
-                  &this_rd_stats.dist, &this_rd_stats.skip, &temp_sse);
+                  &this_rd_stats.dist, &this_rd_stats.skip, &temp_sse, NULL,
+                  NULL, NULL);
   if (av1_is_directional_mode(mbmi->mode) && av1_use_angle_delta(bsize)) {
     mode_cost +=
         x->angle_delta_cost[mbmi->mode - V_PRED]
@@ -6821,7 +6826,7 @@
                                                      this_mode, mi_row, mi_col);
     av1_build_inter_predictors_sby(cm, xd, mi_row, mi_col, ctx, bsize);
     model_rd_for_sb(cpi, bsize, x, xd, 0, 0, &rate_sum, &dist_sum,
-                    &tmp_skip_txfm_sb, &tmp_skip_sse_sb);
+                    &tmp_skip_txfm_sb, &tmp_skip_sse_sb, NULL, NULL, NULL);
     rd = RDCOST(x->rdmult, *rs2 + *out_rate_mv + rate_sum, dist_sum);
     if (rd >= best_rd_cur) {
       mbmi->mv[0].as_int = cur_mv[0].as_int;
@@ -6986,7 +6991,7 @@
   *switchable_rate = av1_get_switchable_rate(cm, x, xd);
   av1_build_inter_predictors_sb(cm, xd, mi_row, mi_col, orig_dst, bsize);
   model_rd_for_sb(cpi, bsize, x, xd, 0, num_planes - 1, &tmp_rate, &tmp_dist,
-                  skip_txfm_sb, skip_sse_sb);
+                  skip_txfm_sb, skip_sse_sb, NULL, NULL, NULL);
   *rd = RDCOST(x->rdmult, *switchable_rate + tmp_rate, tmp_dist);
 
   if (assign_filter == SWITCHABLE) {
@@ -7019,7 +7024,8 @@
           av1_build_inter_predictors_sb(cm, xd, mi_row, mi_col, orig_dst,
                                         bsize);
           model_rd_for_sb(cpi, bsize, x, xd, 0, num_planes - 1, &tmp_rate,
-                          &tmp_dist, &tmp_skip_sb, &tmp_skip_sse);
+                          &tmp_dist, &tmp_skip_sb, &tmp_skip_sse, NULL, NULL,
+                          NULL);
           tmp_rd = RDCOST(x->rdmult, tmp_rs + tmp_rate, tmp_dist);
 
           if (tmp_rd < *rd) {
@@ -7052,7 +7058,8 @@
           av1_build_inter_predictors_sb(cm, xd, mi_row, mi_col, orig_dst,
                                         bsize);
           model_rd_for_sb(cpi, bsize, x, xd, 0, num_planes - 1, &tmp_rate,
-                          &tmp_dist, &tmp_skip_sb, &tmp_skip_sse);
+                          &tmp_dist, &tmp_skip_sb, &tmp_skip_sse, NULL, NULL,
+                          NULL);
           tmp_rd = RDCOST(x->rdmult, tmp_rs + tmp_rate, tmp_dist);
 
           if (tmp_rd < *rd) {
@@ -7086,7 +7093,8 @@
           av1_build_inter_predictors_sb(cm, xd, mi_row, mi_col, orig_dst,
                                         bsize);
           model_rd_for_sb(cpi, bsize, x, xd, 0, num_planes - 1, &tmp_rate,
-                          &tmp_dist, &tmp_skip_sb, &tmp_skip_sse);
+                          &tmp_dist, &tmp_skip_sb, &tmp_skip_sse, NULL, NULL,
+                          NULL);
           tmp_rd = RDCOST(x->rdmult, tmp_rs + tmp_rate, tmp_dist);
 
           if (tmp_rd < *rd) {
@@ -7188,12 +7196,11 @@
       assert(mbmi->ref_frame[1] != INTRA_FRAME);
     }
 
-    // SIMPLE_TRANSLATION mode: no need to recalculate.
-    // The prediction is calculated before motion_mode_rd() is called in
-    // handle_inter_mode()
-
-    // OBMC mode
-    if (mbmi->motion_mode == OBMC_CAUSAL) {
+    if (mbmi->motion_mode == SIMPLE_TRANSLATION && !is_interintra_mode) {
+      // SIMPLE_TRANSLATION mode: no need to recalculate.
+      // The prediction is calculated before motion_mode_rd() is called in
+      // handle_inter_mode()
+    } else if (mbmi->motion_mode == OBMC_CAUSAL) {
       mbmi->motion_mode = OBMC_CAUSAL;
       if (!is_comp_pred && have_newmv_in_inter_mode(this_mode)) {
         int tmp_rate_mv = 0;
@@ -7214,10 +7221,7 @@
       av1_build_obmc_inter_prediction(
           cm, xd, mi_row, mi_col, args->above_pred_buf, args->above_pred_stride,
           args->left_pred_buf, args->left_pred_stride);
-    }
-
-    // Local warped motion mode
-    if (mbmi->motion_mode == WARPED_CAUSAL) {
+    } else if (mbmi->motion_mode == WARPED_CAUSAL) {
       int pts[SAMPLES_ARRAY_SIZE], pts_inref[SAMPLES_ARRAY_SIZE];
       mbmi->motion_mode = WARPED_CAUSAL;
       mbmi->wm_params[0].wmtype = DEFAULT_WMTYPE;
@@ -7279,10 +7283,7 @@
       } else {
         continue;
       }
-    }
-
-    // Interintra mode
-    if (is_interintra_mode) {
+    } else if (is_interintra_mode) {
       INTERINTRA_MODE best_interintra_mode = II_DC_PRED;
       int64_t rd, best_interintra_rd = INT64_MAX;
       int rmode, rate_sum;
@@ -7322,7 +7323,7 @@
                                                   intrapred, bw);
         av1_combine_interintra(xd, bsize, 0, tmp_buf, bw, intrapred, bw);
         model_rd_for_sb(cpi, bsize, x, xd, 0, 0, &rate_sum, &dist_sum,
-                        &tmp_skip_txfm_sb, &tmp_skip_sse_sb);
+                        &tmp_skip_txfm_sb, &tmp_skip_sse_sb, NULL, NULL, NULL);
         rd = RDCOST(x->rdmult, tmp_rate_mv + rate_sum + rmode, dist_sum);
         if (rd < best_interintra_rd) {
           best_interintra_rd = rd;
@@ -7384,7 +7385,8 @@
             av1_build_inter_predictors_sby(cm, xd, mi_row, mi_col, orig_dst,
                                            bsize);
             model_rd_for_sb(cpi, bsize, x, xd, 0, 0, &rate_sum, &dist_sum,
-                            &tmp_skip_txfm_sb, &tmp_skip_sse_sb);
+                            &tmp_skip_txfm_sb, &tmp_skip_sse_sb, NULL, NULL,
+                            NULL);
             rd = RDCOST(x->rdmult, tmp_rate_mv + rmode + rate_sum + rwedge,
                         dist_sum);
             if (rd >= best_interintra_rd_wedge) {
@@ -7740,7 +7742,6 @@
   int refs[2] = { mbmi->ref_frame[0],
                   (mbmi->ref_frame[1] < 0 ? 0 : mbmi->ref_frame[1]) };
   int rate_mv = 0;
-  int pred_exists = 1;
   const int bw = block_size_wide[bsize];
   DECLARE_ALIGNED(32, uint8_t, tmp_buf_[2 * MAX_MB_PLANE * MAX_SB_SQUARE]);
   uint8_t *tmp_buf;
@@ -7787,20 +7788,28 @@
   const RD_STATS backup_rd_stats = *rd_stats;
   const RD_STATS backup_rd_stats_y = *rd_stats_y;
   const RD_STATS backup_rd_stats_uv = *rd_stats_uv;
+  const MB_MODE_INFO backup_mbmi = *mbmi;
+  INTERINTER_COMPOUND_DATA best_compound_data;
+  memset(&best_compound_data, 0, sizeof(best_compound_data));
+  uint8_t tmp_mask_buf[2 * MAX_SB_SQUARE];
+  best_compound_data.seg_mask = tmp_mask_buf;
   RD_STATS best_rd_stats, best_rd_stats_y, best_rd_stats_uv;
   int64_t best_rd = INT64_MAX;
   int best_compound_idx = 1;
   int64_t best_ret_val = INT64_MAX;
   uint8_t best_blk_skip[MAX_MIB_SIZE * MAX_MIB_SIZE];
-  const MB_MODE_INFO backup_mbmi = *mbmi;
   MB_MODE_INFO best_mbmi = *mbmi;
   int64_t early_terminate = 0;
+  int plane_rate[MAX_MB_PLANE] = { 0 };
+  int64_t plane_sse[MAX_MB_PLANE] = { 0 };
+  int64_t plane_dist[MAX_MB_PLANE] = { 0 };
 
   int comp_idx;
   const int search_jnt_comp = is_comp_pred & cm->seq_params.enable_jnt_comp &
                               (mbmi->mode != GLOBAL_GLOBALMV);
   // If !search_jnt_comp, we need to force mbmi->compound_idx = 1.
   for (comp_idx = !search_jnt_comp; comp_idx < 2; ++comp_idx) {
+    rs = 0;
     compmode_interinter_cost = 0;
     early_terminate = 0;
     *rd_stats = backup_rd_stats;
@@ -7825,10 +7834,14 @@
       continue;
     }
     if (have_newmv_in_inter_mode(this_mode)) {
-      ret_val =
-          handle_newmv(cpi, x, bsize, cur_mv, mi_row, mi_col, &rate_mv, args);
+      // when jnt_comp_skip_mv_search flag is on, new mv will be searech once
+      if (!(search_jnt_comp && cpi->sf.jnt_comp_skip_mv_search &&
+            comp_idx == 1))
+        ret_val =
+            handle_newmv(cpi, x, bsize, cur_mv, mi_row, mi_col, &rate_mv, args);
       if (ret_val != 0) {
         early_terminate = INT64_MAX;
+        if (cpi->sf.jnt_comp_skip_mv_search) return INT64_MAX;
         continue;
       } else {
         rd_stats->rate += rate_mv;
@@ -7891,6 +7904,7 @@
         &rd, &rs, &skip_txfm_sb, &skip_sse_sb);
     if (ret_val != 0) {
       early_terminate = INT64_MAX;
+      restore_dst_buf(xd, orig_dst, num_planes);
       continue;
     }
 
@@ -7898,7 +7912,6 @@
       int rate_sum, rs2;
       int64_t dist_sum;
       int64_t best_rd_compound = INT64_MAX, best_rd_cur = INT64_MAX;
-      INTERINTER_COMPOUND_DATA best_compound_data;
       int_mv best_mv[2];
       int best_tmp_rate_mv = rate_mv;
       int tmp_skip_txfm_sb;
@@ -7916,9 +7929,6 @@
 
       best_mv[0].as_int = cur_mv[0].as_int;
       best_mv[1].as_int = cur_mv[1].as_int;
-      memset(&best_compound_data, 0, sizeof(best_compound_data));
-      uint8_t tmp_mask_buf[2 * MAX_SB_SQUARE];
-      best_compound_data.seg_mask = tmp_mask_buf;
 
       if (masked_compound_used) {
         // get inter predictors to use for masked compound modes
@@ -8042,23 +8052,23 @@
         early_terminate = INT64_MAX;
         continue;
       }
-
-      pred_exists = 0;
-
       compmode_interinter_cost = best_compmode_interinter_cost;
     }
 
-    if (pred_exists == 0) {
+    if (is_comp_pred) {
       int tmp_rate;
       int64_t tmp_dist;
       av1_build_inter_predictors_sb(cm, xd, mi_row, mi_col, &orig_dst, bsize);
       model_rd_for_sb(cpi, bsize, x, xd, 0, num_planes - 1, &tmp_rate,
-                      &tmp_dist, &skip_txfm_sb, &skip_sse_sb);
+                      &tmp_dist, &skip_txfm_sb, &skip_sse_sb, plane_rate,
+                      plane_sse, plane_dist);
       rd = RDCOST(x->rdmult, rs + tmp_rate, tmp_dist);
+    }
 
+    if (search_jnt_comp) {
       // if 1/2 model rd is larger than best_rd in jnt_comp mode,
       // use jnt_comp mode, save additional search
-      if ((rd >> 1) > best_rd) {
+      if (comp_idx && (rd >> 1) > best_rd) {
         restore_dst_buf(xd, orig_dst, num_planes);
         continue;
       }
@@ -8096,35 +8106,60 @@
 
     rd_stats->rate += compmode_interinter_cost;
 
-    ret_val = motion_mode_rd(cpi, x, bsize, rd_stats, rd_stats_y, rd_stats_uv,
-                             disable_skip, mi_row, mi_col, args, ref_best_rd,
-                             refs, rate_mv, &orig_dst);
-    if (is_comp_pred && ret_val != INT64_MAX) {
-      int64_t tmp_rd;
-      const int skip_ctx = av1_get_skip_context(xd);
-      if (RDCOST(x->rdmult, rd_stats->rate, rd_stats->dist) <
-          RDCOST(x->rdmult, 0, rd_stats->sse))
-        tmp_rd = RDCOST(x->rdmult, rd_stats->rate + x->skip_cost[skip_ctx][0],
-                        rd_stats->dist);
-      else
-        tmp_rd = RDCOST(x->rdmult,
-                        rd_stats->rate + x->skip_cost[skip_ctx][1] -
-                            rd_stats_y->rate - rd_stats_uv->rate,
-                        rd_stats->sse);
-
-      if (tmp_rd < best_rd) {
-        best_rd_stats = *rd_stats;
-        best_rd_stats_y = *rd_stats_y;
-        best_rd_stats_uv = *rd_stats_uv;
-        best_compound_idx = mbmi->compound_idx;
-        best_ret_val = ret_val;
-        best_rd = tmp_rd;
-        best_mbmi = *mbmi;
-        memcpy(best_blk_skip, x->blk_skip,
-               sizeof(best_blk_skip[0]) * xd->n8_h * xd->n8_w);
+    if (search_jnt_comp) {
+      if (cpi->sf.jnt_comp_fast_tx_search && comp_idx == 0) {
+        // TODO(chengchen): this speed feature introduces big loss.
+        // Need better estimation of rate distortion.
+        rd_stats->rate += rs;
+        rd_stats->rate += plane_rate[0] + plane_rate[1] + plane_rate[2];
+        rd_stats_y->rate = plane_rate[0];
+        rd_stats_uv->rate = plane_rate[1] + plane_rate[2];
+        rd_stats->sse = plane_sse[0] + plane_sse[1] + plane_sse[2];
+        rd_stats_y->sse = plane_sse[0];
+        rd_stats_uv->sse = plane_sse[1] + plane_sse[2];
+        rd_stats->dist = plane_dist[0] + plane_dist[1] + plane_dist[2];
+        rd_stats_y->dist = plane_dist[0];
+        rd_stats_uv->dist = plane_dist[1] + plane_dist[2];
+      } else {
+        ret_val = motion_mode_rd(cpi, x, bsize, rd_stats, rd_stats_y,
+                                 rd_stats_uv, disable_skip, mi_row, mi_col,
+                                 args, ref_best_rd, refs, rate_mv, &orig_dst);
       }
+      if (ret_val != INT64_MAX) {
+        int64_t tmp_rd;
+        const int skip_ctx = av1_get_skip_context(xd);
+        if (RDCOST(x->rdmult, rd_stats->rate, rd_stats->dist) <
+            RDCOST(x->rdmult, 0, rd_stats->sse))
+          tmp_rd = RDCOST(x->rdmult, rd_stats->rate + x->skip_cost[skip_ctx][0],
+                          rd_stats->dist);
+        else
+          tmp_rd = RDCOST(x->rdmult,
+                          rd_stats->rate + x->skip_cost[skip_ctx][1] -
+                              rd_stats_y->rate - rd_stats_uv->rate,
+                          rd_stats->sse);
+
+        if (tmp_rd < best_rd) {
+          best_rd_stats = *rd_stats;
+          best_rd_stats_y = *rd_stats_y;
+          best_rd_stats_uv = *rd_stats_uv;
+          best_compound_idx = mbmi->compound_idx;
+          best_ret_val = ret_val;
+          best_rd = tmp_rd;
+          best_mbmi = *mbmi;
+          memcpy(best_blk_skip, x->blk_skip,
+                 sizeof(best_blk_skip[0]) * xd->n8_h * xd->n8_w);
+        }
+      }
+    } else {
+      ret_val = motion_mode_rd(cpi, x, bsize, rd_stats, rd_stats_y, rd_stats_uv,
+                               disable_skip, mi_row, mi_col, args, ref_best_rd,
+                               refs, rate_mv, &orig_dst);
+      restore_dst_buf(xd, orig_dst, num_planes);
+      if (ret_val != 0) return ret_val;
     }
+    restore_dst_buf(xd, orig_dst, num_planes);
   }
+
   // re-instate status of the best choice
   if (is_comp_pred && best_ret_val != INT64_MAX) {
     *rd_stats = best_rd_stats;
diff --git a/av1/encoder/speed_features.c b/av1/encoder/speed_features.c
index d6c9dc0..e4a84e0 100644
--- a/av1/encoder/speed_features.c
+++ b/av1/encoder/speed_features.c
@@ -434,6 +434,8 @@
   sf->use_inter_txb_hash = 1;
   sf->use_mb_rd_hash = 1;
   sf->optimize_b_precheck = 0;
+  sf->jnt_comp_fast_tx_search = 0;
+  sf->jnt_comp_skip_mv_search = 0;
 
   for (i = 0; i < TX_SIZES; i++) {
     sf->intra_y_mode_mask[i] = INTRA_ALL;
diff --git a/av1/encoder/speed_features.h b/av1/encoder/speed_features.h
index 4e38709..27d68bd 100644
--- a/av1/encoder/speed_features.h
+++ b/av1/encoder/speed_features.h
@@ -576,6 +576,12 @@
 
   // Calculate RD cost before doing optimize_b, and skip if the cost is large.
   int optimize_b_precheck;
+
+  // Use model rd instead of transform search in jnt_comp
+  int jnt_comp_fast_tx_search;
+
+  // Skip mv search in jnt_comp
+  int jnt_comp_skip_mv_search;
 } SPEED_FEATURES;
 
 struct AV1_COMP;