-
Notifications
You must be signed in to change notification settings - Fork 12
Expand file tree
/
Copy pathmodint_montgomery.cpp
More file actions
166 lines (148 loc) · 4.24 KB
/
Copy pathmodint_montgomery.cpp
File metadata and controls
166 lines (148 loc) · 4.24 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
//
// モンゴメリ乗算を用いた modint (mod は 2^62 未満の奇素数)
//
// verified:
// AOJ 2353 Four Arithmetic Operations
// http://judge.u-aizu.ac.jp/onlinejudge/description.jsp?id=2353
//
#include <bits/stdc++.h>
using namespace std;
// montgomery modint (MOD < 2^62, MOD: odd prime number)
struct MontgomeryModInt64 {
using mint = MontgomeryModInt64;
using u64 = uint64_t;
using u128 = __uint128_t;
// static menber
static u64 MOD;
static u64 INV_MOD; // INV_MOD * MOD ≡ 1 (mod 2^64)
static u64 T128; // 2^128 (mod MOD)
// inner value
u64 val;
// constructor
MontgomeryModInt64() : val(0) { }
MontgomeryModInt64(long long v) : val(reduce((u128(v) + MOD) * T128)) { }
u64 get() const {
u64 res = reduce(val);
return res >= MOD ? res - MOD : res;
}
// mod getter and setter
static u64 get_mod() { return MOD; }
static void set_mod(u64 mod) {
assert(mod < (1LL << 62));
assert((mod & 1));
MOD = mod;
T128 = -u128(mod) % mod;
INV_MOD = get_inv_mod();
}
static u64 get_inv_mod() {
u64 res = MOD;
for (int i = 0; i < 5; ++i) res *= 2 - MOD * res;
return res;
}
static u64 reduce(const u128 &v) {
return (v + u128(u64(v) * u64(-INV_MOD)) * MOD) >> 64;
}
// arithmetic operators
mint operator + () const { return mint(*this); }
mint operator - () const { return mint() - mint(*this); }
mint operator + (const mint &r) const { return mint(*this) += r; }
mint operator - (const mint &r) const { return mint(*this) -= r; }
mint operator * (const mint &r) const { return mint(*this) *= r; }
mint operator / (const mint &r) const { return mint(*this) /= r; }
mint& operator += (const mint &r) {
if ((val += r.val) >= 2 * MOD) val -= 2 * MOD;
return *this;
}
mint& operator -= (const mint &r) {
if ((val += 2 * MOD - r.val) >= 2 * MOD) val -= 2 * MOD;
return *this;
}
mint& operator *= (const mint &r) {
val = reduce(u128(val) * r.val);
return *this;
}
mint& operator /= (const mint &r) {
*this *= r.inv();
return *this;
}
mint pow(u128 n) const {
mint res(1), mul(*this);
while (n > 0) {
if (n & 1) res *= mul;
mul *= mul;
n >>= 1;
}
return res;
}
mint inv() const {
return pow(MOD - 2);
}
// other operators
bool operator == (const mint &r) const {
return (val >= MOD ? val - MOD : val) == (r.val >= MOD ? r.val - MOD : r.val);
}
bool operator != (const mint &r) const {
return (val >= MOD ? val - MOD : val) != (r.val >= MOD ? r.val - MOD : r.val);
}
mint& operator ++ () {
++val;
if (val >= MOD) val -= MOD;
return *this;
}
mint& operator -- () {
if (val == 0) val += MOD;
--val;
return *this;
}
mint operator ++ (int) {
mint res = *this;
++*this;
return res;
}
mint operator -- (int) {
mint res = *this;
--*this;
return res;
}
friend istream& operator >> (istream &is, mint &x) {
long long t;
is >> t;
x = mint(t);
return is;
}
friend ostream& operator << (ostream &os, const mint &x) {
return os << x.get();
}
friend mint pow(const mint &r, long long n) {
return r.pow(n);
}
friend mint inv(const mint &r) {
return r.inv();
}
};
typename MontgomeryModInt64::u64
MontgomeryModInt64::MOD, MontgomeryModInt64::INV_MOD, MontgomeryModInt64::T128;
//------------------------------//
// Examples
//------------------------------//
void AOJ2353() {
using mint = MontgomeryModInt64;
const long long MOD = 67280421310721LL;
mint::set_mod(MOD);
mint res = 0;
int N;
cin >> N;
for (int i = 0; i < N; ++i) {
long long o, y;
cin >> o >> y;
if (o == 1) res += y;
else if (o == 2) res -= y;
else if (o == 3) res *= y;
else res /= y;
}
if (res.get() <= 1LL<<31) cout << res.get() << endl;
else cout << (long long)res.get() - MOD << endl;
}
int main() {
AOJ2353();
}