public class Solution {
public int numWays(int n, int k) {
if (n < 1) return n;
if (n == 1) return k;
int same = k;
int diff = k * (k - 1);
for (int i = 2; i < n; i++) {
int temp = same;
same = diff;
diff = (temp + diff) * (k - 1);
}
return same + diff;
}
}