为什么我的Karatsuba比竖式乘法n^2还慢啊
  • 板块学术版
  • 楼主xiaobing
  • 当前回复1
  • 已保存回复1
  • 发布时间2023/8/28 08:42
  • 上次更新2023/11/3 00:47:08
查看原帖
为什么我的Karatsuba比竖式乘法n^2还慢啊
544756
xiaobing楼主2023/8/28 08:42
struct INT {
private:
	vector<int> num;
	int len = 0, ln = 0;
public:
	void add(INT* a, int x, int n) {
		bool b = a->len < x;
		a->len = max(a->len, x);
		while (a->ln < x) {
			a->ln++;
			a->num.push_back(0);
		}
		x--;
		a->num[x] += n;
	}
	INT(string s) {
		int l = s.size();
		if (s[0] == '-')
			for (int i = 1; i < l; i++)
				add(this, i, 0 - (s[l - i] - '0'));
		else
			for (int i = 1; i <= l; i++)
				add(this, i, s[l - i] - '0');
	}
	int size() {
		INT a = *this;
		return a.len;
	}
	long long to_int() {
		INT a = *this;
		long long ans = 0, k = 1;
		for (int i = 0; i < a.len; i++, k *= 10)
			ans += k * a.num[i];
		return ans;
	}
	void print() {
		cout << num[len - 1];
		for (int i = len - 2; i >= 0; i--)
			cout << abs(num[i]);
		return;
	}
	bool operator>(INT b) {
		INT a = *this;
		if (a.num[a.len - 1] == 0) {
			if (b.num[b.len - 1] < 0)
				return 1;
			return 0;
		}
		if (b.num[b.len - 1] == 0) {
			if (a.num[a.len - 1] > 0)
				return 1;
			return 0;
		}
		int ka = a.num[a.len - 1] / abs(a.num[a.len - 1]), kb = b.num[b.len - 1] / abs(b.num[b.len - 1]);
		if (ka == kb) {
			if (a.len != b.len)
				return (a.len > b.len) ^ (ka < 0);
			for (int i = a.len - 1; i >= 0; i--) {
				if (a.num[i] > b.num[i])
					return 1;
				if (a.num[i] < b.num[i])
					return 0;
			}
			return 0;
		}
		else return ka > kb;
	}
	bool operator<(INT b) {
		INT a = *this;
		return b > a;
	}
	bool operator==(INT b) {
		INT a = *this;
		return !((a > b) || (a < b));
	}
	bool operator>=(INT b) {
		INT a = *this;
		return !(a < b);
	}
	bool operator<=(INT b) {
		INT a = *this;
		return !(a > b);
	}
	bool operator!=(INT b) {
		INT a = *this;
		return !(a == b);
	}
	INT _abs() {
		INT ans = *this;
		for (int i = 0; i < len; i++)
			ans.num[i] = abs(num[i]);
		return ans;
	}
	INT subINT(int b, int l) {
		INT ans(""), a = *this;
		for (int i = l - 1; i >= 0; i--)
			add(&ans, i + 1, a.num[a.len - b - l + i + 1]);
		while (ans.num[ans.len - 1] == 0 && ans.len > 1)
			ans.len--;
		return ans;
	}
	INT operator+(INT b) {
		INT a = *this, zero("0");
		if (a.num[a.len - 1] == 0)
			return b;
		if (b.num[b.len - 1] == 0)
			return a;
		int ka = a.num[a.len - 1] / abs(a.num[a.len - 1]), kb = b.num[b.len - 1] / abs(b.num[b.len - 1]);
		if (ka != kb) {
			if (a < b)
				swap(a, b);
			return a - (zero - b);
		}
		if (ka < 0)
			return zero - (a._abs() + b._abs());
		if (a.len < b.len)
			swap(a, b);
		INT ans = a;
		for (int i = 0; i < b.len; i++)
			add(&ans, i + 1, b.num[i]);
		for (int i = 0; i < a.len; i++)
			if (ans.num[i] >= 10) {
				add(&ans, i + 2, ans.num[i] / 10);
				ans.num[i] %= 10;
			}
		while (ans.num[ans.len - 1] == 0)
			ans.len--;
		return ans;
	}
	void operator+=(INT b) {
		INT a = *this;
		*this = a + b;
		return;
	}
	INT operator-(INT b) {
		INT a = *this, zero("0"), ans("");
		if (a == zero) {
			INT ans = b;
			for (int i = 0; i < ans.len; i++)
				ans.num[i] = -ans.num[i];
			return ans;
		}
		if (b.num[b.len - 1] == 0)
			return a;
		int ka = a.num[a.len - 1] / abs(a.num[a.len - 1]), kb = b.num[b.len - 1] / abs(b.num[b.len - 1]);
		if (kb < 0)
			return a + b._abs();
		if (a == b)
			return zero;
		if (ka != kb)
			return zero - (a._abs() + b._abs());
		if (a < b)
			return zero - (b - a);
		ans = a;
		for (int i = 0; i < b.len; i++)
			add(&ans, i + 1, -b.num[i]);
		for (int i = 0; i < ans.len; i++)
			if (ans.num[i] < 0) {
				add(&ans, i + 1, 10);
				add(&ans, i + 2, -1);
			}
		while (!ans.num[ans.len - 1])
			ans.len--;
		return ans;
	}
	void operator-=(INT b) {
		INT a = *this;
		*this = a - b;
		return;
	}
	INT Karatsuba(INT a, INT b, int n) {
		if (a.len < b.len)
			swap(a, b);
		if (n < 9) {
			long long A = a.to_int(), B = b.to_int();
			long long ans = 1;
			ans = A * B;
			string s = to_string(ans);
			INT rt(s);
			return rt;
		}
		INT ans(""), zero("0"), z1(""), z2(""), z3("");
		int k = n >> 1;
		INT a1 = a.subINT(1, n - k), a2 = a.subINT(n - k + 1, k), b1(""), b2("");
		if (b.len > k) {
			b1 = b.subINT(1, b.len - k);
			b2 = b.subINT(b.len - k + 1, k);
			z1 = Karatsuba(a1, b1, max(a1.len, b1.len));
			z2 = Karatsuba(a1, b2, max(a1.len, b2.len)) + Karatsuba(a2, b1, max(a2.len, b1.len));
			z3 = Karatsuba(a2, b2, max(a2.len, b2.len));
		}
		else {
			z1 = zero;
			z2 = Karatsuba(a1, b, max(a1.len, b.len));
			z3 = Karatsuba(a2, b, max(a2.len, b.len));
		}
		for (int i = 0; i < z1.len; i++)
			add(&ans, i + 1 + 2 * k, z1.num[i]);
		for (int i = 0; i < z2.len; i++)
			add(&ans, i + 1 + k, z2.num[i]);
		for (int i = 0; i < z3.len; i++)
			add(&ans, i + 1, z3.num[i]);
		for (int i = 0; i < ans.len; i++)
			if (ans.num[i] >= 10) {
				add(&ans, i + 2, ans.num[i] / 10);
				ans.num[i] %= 10;
			}
		while (ans.num[ans.len - 1] == 0)
			ans.len--;
		return ans;
	}
	INT operator*(INT b) {
		INT a = *this, zero("0"), ans("");
		if (a.num[a.len - 1] == 0 || b.num[b.len - 1] == 0)
			return zero;
		int ka = a.num[a.len - 1] / abs(a.num[a.len - 1]), kb = b.num[b.len - 1] / abs(b.num[b.len - 1]);
		if (ka != kb)
			return zero - ((zero - a) * b);
		if (a < b)
			swap(a, b);
//		for (int i = 0; i < a.len; i++)
//			for (int j = 0; j < b.len; j++)
//				add(&ans, i + j + 1, a.num[i] * b.num[j]);
//		for (int i = 0; i < ans.len; i++)
//			if (ans.num[i] >= 10) {
//				add(&ans, i + 2, ans.num[i] / 10);
//				ans.num[i] %= 10;
//			}
		ans = Karatsuba(a, b, max(a.len, b.len));
		return ans;
	}
	void operator*=(INT b) {
		INT a = *this;
		*this = a * b;
		return;
	}
	INT operator/(INT b) {
		INT a = *this, zero("0");
		if (a == zero)
			return zero;
		int ka = a.num[a.len - 1] / abs(a.num[a.len - 1]), kb = b.num[b.len - 1] / abs(b.num[b.len - 1]);
		if (ka != kb) {
			INT ans = zero - (a._abs() / b._abs()), one("1");
			if (ans * b != a)
				ans -= one;
			return ans;
		}
		if (a < b)
			return zero;
		int la = b.len;
		string s;
		INT cnt("0"), ten("10");
		for (int i = a.len - 1; i >= 0; i--) {
			INT ka(to_string(a.num[i]));
			cnt = cnt * ten + ka;
			int l = 0, r = 10;
			if (cnt >= b)
				while (l + 1 < r) {
					int mid = (l + r) >> 1;
					INT k(to_string(mid));
					if (k * b > cnt)
						r = mid;
					else l = mid;
				}
			if (l != 0 || !s.empty())
				s += l + '0';
			INT k(to_string(l));
			k *= b;
			cnt -= k;
		}
		INT ans(s);
		while (ans.num[ans.len - 1] == 0)
			ans.len--;
		return ans;
	}
	void operator/=(INT b) {
		INT a = *this;
		*this = a / b;
		return;
	}
	INT operator%(INT b) {
		INT a = *this;
		return a - (a / b * b);
	}
	void operator%=(INT b) {
		INT a = *this;
		*this = a % b;
		return;
	}
	INT pow(int k) {
		INT one("1"), n = *this;
		if (k == 0)
			return one;
		if (k == 1)
			return n;
		INT sum = n.pow(k / 2);
		sum *= sum;
		if (k % 2)
			sum *= n;
		return sum;
	}
	INT root(int k) {
		INT a = *this, zero("0"), one("1"), two("2"), ten("10");
		if (a == zero)
			return zero;
		int l = 0, r = a.len;
		while (l + 1 < r) {
			int mid = (l + r) >> 1;
			if (mid * k + 1 > a.len)
				r = mid;
			else l = mid;
		}
		INT L = ten.pow(l), R = ten.pow(r);
		while ((L + one) < R) {
			INT mid = (L + R) / two;
			if (mid.pow(k) > a)
				R = mid;
			else L = mid;
		}
		return L;
	}
};
2023/8/28 08:42
加载中...