AtCoderのABC    前のABCの問題へ

ABC471-E Sum of Square of Sum


問題へのリンク


C#のソース

using System;
using System.Collections.Generic;
using System.Linq;

class Program
{
    static string InputPattern = "InputX";

    static List<string> GetInputList()
    {
        var WillReturn = new List<string>();

        if (InputPattern == "Input1") {
            WillReturn.Add("3 2");
            WillReturn.Add("1 10 100");
            //22422
        }
        else if (InputPattern == "Input2") {
            WillReturn.Add("5 2");
            WillReturn.Add("10 10 20 20 20");
            //10600
        }
        else if (InputPattern == "Input3") {
            WillReturn.Add("2 1");
            WillReturn.Add("998244353 998244353");
            //0
        }
        else {
            string wkStr;
            while ((wkStr = Console.ReadLine()) != null) WillReturn.Add(wkStr);
        }
        return WillReturn;
    }

    static long[] GetSplitArr(string pStr)
    {
        return (pStr == "" ? new string[0] : pStr.Split(' ')).Select(pX => long.Parse(pX)).ToArray();
    }

    const long Hou = 998244353;

    static void Main()
    {
        List<string> InputList = GetInputList();
        long[] wkArr = GetSplitArr(InputList[0]);
        long N = wkArr[0];
        long K = wkArr[1];

        long[] AArr = GetSplitArr(InputList[1]);
        long UB = AArr.GetUpperBound(0);

        long Answer = 0;

        var InsChooseMod = new ChooseMod(N, Hou);
        var Ins_Fenwick_Tree = new Fenwick_Tree(AArr, Hou);

        if (K == 1) {
            foreach (long EachA in AArr) {
                long SquareVal = EachA * EachA;
                SquareVal %= Hou;
                Answer += SquareVal;
                Answer %= Hou;
            }
            Console.WriteLine(Answer);
            return;
        }
        if (K == 2) {
            foreach (long EachA in AArr) {
                long SquareVal = EachA * EachA;
                SquareVal %= Hou;
                Answer += SquareVal * (N - 1);
                Answer %= Hou;
            }

            for (long I = 0; I <= UB; I++) {
                long RangeSta = I + 1;
                long RangeEnd = UB;
                if (RangeSta <= RangeEnd) {
                    long RangeSum = Ins_Fenwick_Tree.GetSum(RangeSta, RangeEnd);
                    RangeSum *= 2 * AArr[I];
                    RangeSum %= Hou;
                    Answer += RangeSum;
                    Answer %= Hou;
                }
            }
            Console.WriteLine(Answer);
            return;
        }
        foreach (long EachA in AArr) {
            long SquareVal = EachA * EachA;
            SquareVal %= Hou;
            Answer += SquareVal * InsChooseMod.DeriveChoose(N - 1, K - 1);
            Answer %= Hou;
        }

        for (long I = 0; I <= UB; I++) {
            long RangeSta = I + 1;
            long RangeEnd = UB;
            if (RangeSta <= RangeEnd) {
                long RangeSum = Ins_Fenwick_Tree.GetSum(RangeSta, RangeEnd);
                RangeSum *= 2 * AArr[I];
                RangeSum %= Hou;
                RangeSum *= InsChooseMod.DeriveChoose(N - 2, K - 2);
                RangeSum %= Hou;
                Answer += RangeSum;
                Answer %= Hou;
            }
        }
        Console.WriteLine(Answer);
    }
}

#region ChooseMod
// 二項係数クラス (nCr を nの最大値指定で事前準備し、高速に求める)
internal class ChooseMod
{
    private long mHou;

    private long[] mFacArr;
    private long[] mFacInvArr;
    private long[] mInvArr;

    // コンストラクタ
    internal ChooseMod(long pCnt, long pHou)
    {
        mHou = pHou;
        mFacArr = new long[pCnt + 1];
        mFacInvArr = new long[pCnt + 1];
        mInvArr = new long[pCnt + 1];

        mFacArr[0] = mFacArr[1] = 1;
        mFacInvArr[0] = mFacInvArr[1] = 1;
        mInvArr[1] = 1;
        for (int I = 2; I <= pCnt; I++) {
            mFacArr[I] = mFacArr[I - 1] * I % mHou;
            mInvArr[I] = mHou - mInvArr[mHou % I] * (mHou / I) % mHou;
            mFacInvArr[I] = mFacInvArr[I - 1] * mInvArr[I] % mHou;
        }
    }

    // nCrを返す
    internal long DeriveChoose(long pN, long pR)
    {
        if (pN < pR) return 0;
        if (pN < 0 || pR < 0) return 0;
        return mFacArr[pN] * (mFacInvArr[pR] * mFacInvArr[pN - pR] % mHou) % mHou;
    }
}
#endregion

// フェニック木
#region Fenwick_Tree
internal class Fenwick_Tree
{
    private long[] mBitArr;
    private long mExternalArrUB;
    private long mHou;

    // ノードのIndexの列挙を返す
    internal IEnumerable<long> GetNodeIndEnum()
    {
        for (long I = 0; I <= mExternalArrUB; I++) {
            yield return I;
        }
    }

    // 木のノードのUBを返す
    internal long GetUB()
    {
        return mExternalArrUB;
    }

    // コンストラクタ(外部配列のUBと法を指定)
    internal Fenwick_Tree(long pExternalArrUB, long pHou)
    {
        mExternalArrUB = pExternalArrUB;

        // フェニック木の外部配列は0オリジンで、
        // フェニック木の内部配列は1オリジンなため、2を足す
        mBitArr = new long[pExternalArrUB + 2];

        mHou = pHou;
    }

    // コンストラクタ(初期化用の配列と法を指定)
    internal Fenwick_Tree(long[] pArr, long pHou)
        : this(pArr.GetUpperBound(0), pHou)
    {
        for (long I = 0; I <= pArr.GetUpperBound(0); I++) {
            this.Add(I, pArr[I]);
        }
    }

    // コンストラクタ(初期化用のListと法を指定)
    internal Fenwick_Tree(List<long> pList, long pHou)
        : this(pList.Count - 1, pHou)
    {
        for (int I = 0; I <= pList.Count - 1; I++) {
            this.Add(I, pList[I]);
        }
    }

    // Indのチェック
    private void IndCheck(long pInd)
    {
        if (pInd < 0) throw new Exception("pInd < 0");
        if (mExternalArrUB < pInd) throw new Exception("UB < pInd");
    }

    // Indの大小チェック
    private void IndRangeCheck(long pSta, long pEnd)
    {
        IndCheck(pSta);
        IndCheck(pEnd);
        if (pSta > pEnd) throw new Exception("pSta > pEnd");
    }

    // インデクサ
    internal long this[long pInd]
    {
        get { return GetSum(pInd, pInd); }
        set { Add(pInd, value - GetSum(pInd, pInd)); }
    }

    // [pSta,pEnd] のSumを返す
    internal long GetSum(long pSta, long pEnd)
    {
        IndRangeCheck(pSta, pEnd);

        long Result = GetSum(pEnd);
        if (pSta > 0) {
            Result -= GetSum(pSta - 1);
        }

        Result %= mHou;
        if (Result < 0) Result += mHou;
        return Result;
    }

    // [0,pEnd] のSumを返す
    internal long GetSum(long pEnd)
    {
        IndCheck(pEnd);

        pEnd++; // 1オリジンに変更

        long Sum = 0;
        while (pEnd >= 1) {
            Sum += mBitArr[pEnd];
            Sum %= mHou;
            pEnd -= pEnd & -pEnd;
        }
        if (Sum < 0) Sum += mHou;
        return Sum;
    }

    // [I] に Xを加算
    internal void Add(long pI, long pX)
    {
        IndCheck(pI);

        pI++; // 1オリジンに変更

        pX %= mHou;
        while (pI <= mBitArr.GetUpperBound(0)) {
            mBitArr[pI] += pX;
            mBitArr[pI] %= mHou;
            pI += pI & -pI;
        }
    }
}
#endregion


解説

A B C D
として、
Kを場合分けして考えます。

K=1の場合、平方数の総和であることが自明です。
K=2の場合、図を書いて考えます。
   A    B    C    D
A  A^2  AB   AC   AD
B  AB   B^2  BC   BD
C  AC   BC   C^2  CD
D  AD   BD   CD   D^2

A^2の寄与度は、A以外とのペア数なので、A^2 * 3 です。
ABの寄与度は、AB * 2です。
ACの寄与度は、AB * 2です。
ADの寄与度は、AD * 2です。
これらは、2A(B+C+D)と分配法則でまとめることができ、
B+C+Dは、フェニック木で区間和で高速に求めることができます。

Kが3以上の場合を考えます。
   A    B    C    D
A  A^2  AB   AC   AD
B  AB   B^2  BC   BD
C  AC   BC   C^2  CD
D  AD   BD   CD   D^2
K=2の場合と同様に寄与度を考えます。
A^2の寄与度は、A^2 * (Aとその他からなる組合せ数)です。
ABの寄与度は、2AB * (AとBとその他からなる組合せ数)です。
ACの寄与度は、2AC * (AとCとその他からなる組合せ数)です。
ADの寄与度は、2AD * (AとDとその他からなる組合せ数)です。
これらは、2A(B+C+D)と分配法則でまとめることができ、
B+C+Dは、フェニック木で区間和で高速に求めることができます。

以上の考察により、解くことができます。