Use macro to set txk_type

This will make txk_sel support maximum bsize to 128x128

Change-Id: I33941966cb1ae4406ac68a2124c859c833a084d8
diff --git a/av1/common/blockd.h b/av1/common/blockd.h
index afc92f4..eeff0f2 100644
--- a/av1/common/blockd.h
+++ b/av1/common/blockd.h
@@ -1021,12 +1021,18 @@
   if (xd->lossless[mbmi->segment_id] || txsize_sqr_map[tx_size] >= TX_32X32) {
     tx_type = DCT_DCT;
   } else {
-    if (plane_type == PLANE_TYPE_Y)
-      tx_type = mbmi->txk_type[(blk_row << 4) + blk_col];
-    else if (is_inter_block(mbmi))
-      tx_type = mbmi->txk_type[(blk_row << 5) + (blk_col << 1)];
-    else
+    if (plane_type == PLANE_TYPE_Y) {
+      tx_type = mbmi->txk_type[(blk_row << MAX_MIB_SIZE_LOG2) + blk_col];
+    } else if (is_inter_block(mbmi)) {
+      // scale back to y plane's coordinate
+      blk_row <<= pd->subsampling_y;
+      blk_col <<= pd->subsampling_x;
+      tx_type = mbmi->txk_type[(blk_row << MAX_MIB_SIZE_LOG2) + blk_col];
+    } else {
+      // In intra mode, uv planes don't share the same prediction mode as y
+      // plane, so the tx_type should not be shared
       tx_type = intra_mode_to_tx_type_context[mbmi->uv_mode];
+    }
   }
   assert(tx_type >= DCT_DCT && tx_type < TX_TYPES);
   if (is_inter_block(mbmi) && !av1_ext_tx_used[tx_set_type][tx_type])
diff --git a/av1/decoder/decodemv.c b/av1/decoder/decodemv.c
index 677dfae..085ddd9 100644
--- a/av1/decoder/decodemv.c
+++ b/av1/decoder/decodemv.c
@@ -933,7 +933,7 @@
   // only y plane's tx_type is transmitted
   if (plane > 0) return;
   (void)block;
-  TX_TYPE *tx_type = &mbmi->txk_type[(blk_row << 4) + blk_col];
+  TX_TYPE *tx_type = &mbmi->txk_type[(blk_row << MAX_MIB_SIZE_LOG2) + blk_col];
 #endif
 
   if (!FIXED_TX_TYPE) {
diff --git a/av1/decoder/decodetxb.c b/av1/decoder/decodetxb.c
index 3f19372..56619f9 100644
--- a/av1/decoder/decodetxb.c
+++ b/av1/decoder/decodetxb.c
@@ -95,7 +95,8 @@
   if (all_zero) {
     *max_scan_line = 0;
 #if CONFIG_TXK_SEL
-    if (plane == 0) mbmi->txk_type[(blk_row << 4) + blk_col] = DCT_DCT;
+    if (plane == 0)
+      mbmi->txk_type[(blk_row << MAX_MIB_SIZE_LOG2) + blk_col] = DCT_DCT;
 #endif
     return 0;
   }
diff --git a/av1/encoder/encodetxb.c b/av1/encoder/encodetxb.c
index 50a5a99..7094e83 100644
--- a/av1/encoder/encodetxb.c
+++ b/av1/encoder/encodetxb.c
@@ -2514,7 +2514,8 @@
   av1_invalid_rd_stats(&best_rd_stats);
 
   for (tx_type = txk_start; tx_type <= txk_end; ++tx_type) {
-    if (plane == 0) mbmi->txk_type[(blk_row << 4) + blk_col] = tx_type;
+    if (plane == 0)
+      mbmi->txk_type[(blk_row << MAX_MIB_SIZE_LOG2) + blk_col] = tx_type;
     TX_TYPE ref_tx_type = av1_get_tx_type(get_plane_type(plane), xd, blk_row,
                                           blk_col, block, tx_size);
     if (tx_type != ref_tx_type) {
@@ -2555,7 +2556,8 @@
 
   if (best_eob == 0 && is_inter_block(mbmi)) best_tx_type = DCT_DCT;
 
-  if (plane == 0) mbmi->txk_type[(blk_row << 4) + blk_col] = best_tx_type;
+  if (plane == 0)
+    mbmi->txk_type[(blk_row << MAX_MIB_SIZE_LOG2) + blk_col] = best_tx_type;
   x->plane[plane].txb_entropy_ctx[block] = best_eob;
 
   if (!is_inter_block(mbmi)) {
diff --git a/av1/encoder/rdopt.c b/av1/encoder/rdopt.c
index 84aa8fd..58b789f 100644
--- a/av1/encoder/rdopt.c
+++ b/av1/encoder/rdopt.c
@@ -2695,7 +2695,9 @@
       ref_best_rd = AOMMIN(rd, ref_best_rd);
       if (rd < best_rd) {
 #if CONFIG_TXK_SEL
-        memcpy(best_txk_type, mbmi->txk_type, sizeof(best_txk_type[0]) * 256);
+        memcpy(best_txk_type, mbmi->txk_type,
+               sizeof(best_txk_type[0]) * MAX_SB_SQUARE /
+                   (TX_SIZE_W_MIN * TX_SIZE_H_MIN));
 #endif
         best_tx_type = tx_type;
         best_tx_size = n;
@@ -2711,7 +2713,9 @@
   mbmi->tx_size = best_tx_size;
   mbmi->tx_type = best_tx_type;
 #if CONFIG_TXK_SEL
-  memcpy(mbmi->txk_type, best_txk_type, sizeof(best_txk_type[0]) * 256);
+  memcpy(mbmi->txk_type, best_txk_type,
+         sizeof(best_txk_type[0]) * MAX_SB_SQUARE /
+             (TX_SIZE_W_MIN * TX_SIZE_H_MIN));
 #endif
 
   mbmi->min_tx_size = get_min_tx_size(mbmi->tx_size);
@@ -3857,7 +3861,7 @@
   RD_STATS sum_rd_stats;
 #if CONFIG_TXK_SEL
   TX_TYPE best_tx_type = TX_TYPES;
-  int txk_idx = (blk_row << 4) + blk_col;
+  int txk_idx = (blk_row << MAX_MIB_SIZE_LOG2) + blk_col;
 #endif
 
   av1_init_rd_stats(&sum_rd_stats);