3 条题解

  • 0
    @ 2026-2-13 14:54:54
    #include<bits/stdc++.h>
    using namespace std;
    #define int long long
    #define ac (int)floor(1.0*a/c)
    #define bc (int)floor(1.0*b/c)
    #define gs(n) n*(n+1)/2
    int f(int a,int b,int c,int n)
    {
    	if(a==0)return bc*(n+1);
    	if(n==0)return bc;
    	if(a>=c||b>=c)
    	{
    		int sum=f(a%c,b%c,c,n);
    		return gs(n)*ac+(n+1)*bc+sum;
    	}
    	int m=floor(1.0*(a*n+b)/c),sum=f(c,c-b-1,a,m-1);
    	return n*m-sum;
    }
    signed main()
    {
    	int t;cin>>t;
    	while(t--)
    	{
    		int n,m,a,b;cin>>n>>m>>a>>b;
    		if(a==0)
    		{
    			cout<<min(n,b)<<'\n';
    			continue;
    		}
    		int sum1=f(a-1,b+m-1,m,n-1),sum2=f(a,b,m,n-1);
    		cout<<sum1-sum2<<'\n';
    	}
    	return 0;
    }
    • 0
      @ 2026-2-12 22:36:54

      #include<bits/stdc++.h>
      using namespace std;
      namespace atcoder
      {
      	using Int = long long;
      
      	Int floor_div(Int a, Int b)
      	{
      		Int q = a / b;
      		Int r = a % b;
      		if (r != 0 && a < 0)
      			q -= 1;
      		return q;
      	}
      
      	Int floor_sum_nonneg(Int n, Int m, Int a, Int b)
      	{
      		Int ans = 0;
      		while (true)
      		{
      			if (a >= m)
      			{
      				ans += (n - 1) * n / 2 * (a / m);
      				a %= m;
      			}
      			if (b >= m)
      			{
      				ans += n * (b / m);
      				b %= m;
      			}
      			Int y_max = (a * n + b) / m;
      			if (y_max == 0)
      				return ans;
      			Int x_max = y_max * m - b;
      			ans += (n - (x_max + a - 1) / a) * y_max;
      			n = y_max;
      			b = (a - x_max % a) % a;
      			std::swap(a, m);
      		}
      	}
      
      	Int floor_sum(Int n, Int m, Int a, Int b)
      	{
      		Int ans = 0;
      		Int a_div = atcoder::floor_div(a, m);
      		Int b_div = atcoder::floor_div(b, m);
      		Int a_mod = a - a_div * m;
      		Int b_mod = b - b_div * m;
      		ans += a_div * n * (n - 1) / 2;
      		ans += b_div * n;
      		ans += atcoder::floor_sum_nonneg(n, m, a_mod, b_mod);
      		return ans;
      	}
      
      	void main()
      	{
      		Int T;
      		std::cin >> T;
      		for (Int t = 0; t < T; t++)
      		{
      			Int N, M, A, B;
      			std::cin >> N >> M >> A >> B;
      			Int total = N - atcoder::floor_sum(N, M, A, B) + atcoder::floor_sum(N, M, A - 1, B - 1);
      			std::cout << total << std::endl;
      		}
      	}
      }
      
      int main()
      {
      	atcoder::main();
      	return 0;
      }
      
      
      • 0
        @ 2026-2-12 22:26:07

        abc443-g 题解

        前言

        为什么大家都用 atcoder 自带的写,这显得我是小丑诶。

        终于 AK 了。

        思路

        转化一下:

        $$\begin{aligned} &\phantom{\iff} k < (Ak+B) \bmod M\\ &\iff k+1 \le Ak+B-M\left\lfloor\frac{Ak+B}M \right\rfloor \\ &\iff \left\lfloor\frac{Ak+B}M \right\rfloor \le \frac{(A-1)k+(B-1)}M\\ &\iff \left\lfloor\frac{Ak+B}M \right\rfloor \le \left\lfloor\frac{(A-1)k+(B-1)}M\right\rfloor\\ \end{aligned}$$

        注意到 $\displaystyle \left\lfloor\frac{Ak+B}M \right\rfloor-\left\lfloor\frac{(A-1)k+(B-1)}M \right\rfloor$ 的值只可能是 0 或 1。

        只要求 $\displaystyle N-\sum_{k=0}^{N-1} \left(\left\lfloor\frac{Ak+B}M \right\rfloor-\left\lfloor\frac{(A-1)k+(B-1)}M \right\rfloor\right)$ 即可。

        运用 floor sum 算法即可优化至对数时间复杂度。

        code

        #include <bits/stdc++.h>
        using namespace std;
        
        #define ll __int128
        
        ll exgcd(ll a, ll b, ll &x, ll &y)
        {
            if (b == 0)
            {
                x = (a >= 0 ? 1 : -1);
                y = 0;
                return llabs(a);
            }
        
            ll x1, y1;
            ll g = exgcd(b, a % b, x1, y1);
            x = y1;
            y = x1 - (a / b) * y1;
            return g;
        }
        
        ll getsum(ll n, ll m, ll a, ll b)
        {
            ll ans = 0;
        
            while (1)
            {
                if (a >= m)
                {
                    ll q = a / m;
                    ans += (n - 1) * n / 2 * q;
                    a %= m;
                }
        
                if (b >= m)
                {
                    ll q = b / m;
                    ans += n * q;
                    b %= m;
                }
        
                ll y = a * n + b;
                if (y < m)
                    break;
        
                n = y / m;
                b = y % m;
                swap(a, m);
            }
        
            return ans;
        }
        
        void please_ac()
        {
            long long n, m, a, b;
            cin >> n >> m >> a >> b;
        
            if (a == 0)
            {
                ll r = b % m;
                long long ans = min((ll)n, r);
        
                cout << ans << "\n";
                return;
            }
        
            ll c = a - 1;
            ll sa = getsum(n, m, a, b), sc = getsum(n, m, c, b), s2 = sa - sc;
        
            ll cnt0 = 0;
            if (c == 0)
            {
                if ((b % m) == 0)
                    cnt0 = n;
                else
                    cnt0 = 0;
            }
            else
            {
                ll x, y;
                ll g = __gcd(llabs(c), m);
        
                if ((b % g) != 0)
                    cnt0 = 0;
                else
                {
                    ll c_ = c / g;
                    ll m_ = m / g;
        
                    ll rhs = ((-b / g) % m_ + m_) % m_;
        
                    ll x_, y_;
                    ll gc = exgcd(c_, m_, x_, y_);
        
                    ll inv = (x_ % m_ + m_) % m_;
                    ll k0 = ((__int128)rhs * inv) % m_;
        
                    if (k0 < n)
                        cnt0 = 1 + (n - 1 - k0) / m_;
                    else
                        cnt0 = 0;
                }
            }
        
            ll s1 = n - cnt0;
        
            long long ans = max((ll)0, s1 - s2);
            cout << ans << "\n";
        }
        
        int main()
        {
            ios::sync_with_stdio(0);
            cin.tie(0), cout.tie(0);
        
            int T_T = 1;
            cin >> T_T;
            while (T_T--)
                please_ac();
        }
        
        • 1

        信息

        ID
        973
        时间
        2000ms
        内存
        1024MiB
        难度
        9
        标签
        递交数
        22
        已通过
        4
        上传者