题解:P17212 [ICPC 2017 Nanning R] Attacker-Defender Game

· · 题解

题意简述

进攻方与防守方每回合同时选择消耗的能量。比较结果决定防守方生命值的增减,并扣除相应能量。求双方采用最优混合策略时,进攻方的获胜概率。

解题思路

f_{h,d,a} 表示当前状态下进攻方的最优获胜概率。双方仍有能量时,每回合选择正整数;能量耗尽后只能选择 0。边界为:

f_{h,d,a}= \begin{cases} 1 & h=0 \\ 0 & a<h \\ 1 & d=0\land a\ge h \end{cases}

只讨论 h,d\ge1a\ge h 的状态。进攻方选择 1\le x\le a,防守方选择 1\le y\le d,这一回合后的收益矩阵为:

G_{x,y}= \begin{cases} f_{h-1,d,a-x} & x>y \\ f_{h+1,d-y,a} & x<y \\ f_{h,d-y,a-x} & x=y \end{cases}

每个后继状态的 a+d 都严格减小,因此按 a+d 递增计算即可。

先排除取值为 01 的状态。令 r_{h,d} 为使获胜概率为正所需的最少进攻方能量。边界是 r_{0,d}=0r_{h,0}=h

固定防守方本回合选择的 y。进攻方可以选择 x=y+1 造成伤害,选择 x=y,或在 y\ge2 时选择 x=1<y。三类选择所需的最少能量分别为:

\begin{aligned} & y+1+r_{h-1,d} \\ & y+r_{h,d-y} \\ & r_{h+1,d-y} \end{aligned}

若每一列至少有一个正元素,进攻方给所有行动分配正概率,就能保证期望收益为正;若某列全为零,防守方固定选择该列即可。因此:

r_{h,d}=\max_{1\le y\le d} \begin{cases} \min\left(y+1+r_{h-1,d},y+r_{h,d-y}\right) & y=1 \\ \min\left(y+1+r_{h-1,d},y+r_{h,d-y},r_{h+1,d-y}\right) & y\ge2 \end{cases}

a<r_{h,d} 时,答案恰为 0

答案恰为 1 当且仅当 a\ge h+d。充分性可对 d 归纳。进攻方固定选择 x=1,任意后继状态仍满足对应不等式。

必要性可对 a+d 归纳。若 a<h+d,对进攻方的任意选择 x 分两种情况:当 x\le d 时,防守方取 y=x,后继仍有 a-x<h+d-x;当 x>d 时,防守方取 y=d,此时必有 h\ge2,且 a-x<(h-1)+d。归纳假设说明相应后继都不是必胜态,所以收益矩阵不存在全为 1 的行。

只对满足:

r_{h,d}\le a<h+d

的状态求解矩阵博弈。

剩余状态通过线性规划求解矩阵博弈。令 B=G+1。设防守方变量为 z_y,求解:

\begin{aligned} \max & \sum_{y=1}^{d}z_y \\ \text{s.t.} & \sum_{y=1}^{d}B_{x,y}z_y\le1 & (1\le x\le a) \\ & z_y\ge0 \end{aligned}

若最优目标值为 Z,将 z 除以 Z 就得到防守方的概率分布。它能把矩阵 B 的期望收益限制在 1/Z。由线性规划对偶,进攻方也能保证该值,故:

f_{h,d,a}=\frac{1}{Z}-1

约束右端均为 1,原点可以作为初始基本可行解。代码使用单纯形法求 Z,并用 Bland 规则避免退化时循环。

阈值预处理为 O(AD^2)。对每个确需处理的 a\times d 状态,构造矩阵与单次换基均为 O(ad)。单纯形法理论最坏为指数级,但本题中 A,D<50,且两类阈值会跳过大量状态。空间复杂度为 O(A^2D+AD)

参考代码

#include <bits/stdc++.h>
using namespace std;

using ld=long double;
const int N=55;
const int Q=1205;
const ld eps=1e-15L;
ld f[N][N][N],g[N][N],p[N][N];
int need[N][N],q[Q][3],bs[N],nb[N];
void pivot(int r,int s,int m,int n)
{
    ld inv=1/p[r][s];
    for(int i=0;i<=m;i++)
    {
        if(i==r)continue;
        for(int j=0;j<=n;j++)
        {
            if(j==s)continue;
            p[i][j]-=p[r][j]*p[i][s]*inv;
        }
    }
    for(int i=0;i<=n;i++)if(i!=s)p[r][i]*=inv;
    for(int i=0;i<=m;i++)if(i!=r)p[i][s]*=-inv;
    p[r][s]=inv;
    swap(bs[r],nb[s]);
}
ld simplex(int m,int n)
{
    for(int i=0;i<m;i++)
    {
        bs[i]=n+i;
        p[i][n]=1;
        for(int j=0;j<n;j++)p[i][j]=g[i][j]+1;
    }
    for(int i=0;i<n;i++)
    {
        nb[i]=i;
        p[m][i]=-1;
    }
    p[m][n]=0;
    while(1)
    {
        int s=-1;
        for(int i=0;i<n;i++)
        {
            if(p[m][i]<-eps&&(s==-1||nb[i]<nb[s]))s=i;
        }
        if(s==-1)return p[m][n];
        int r=-1;
        for(int i=0;i<m;i++)
        {
            if(p[i][s]<=eps)continue;
            if(r==-1)
            {
                r=i;
                continue;
            }
            ld x=p[i][n]/p[i][s],y=p[r][n]/p[r][s];
            if(x<y-eps||(fabsl(x-y)<=eps&&bs[i]<bs[r]))r=i;
        }
        if(r==-1)return p[m][n];
        pivot(r,s,m,n);
    }
}
ld get(int h,int d,int a)
{
    if(!h)return 1;
    if(a<h)return 0;
    if(!d)return 1;
    return f[h][d][a];
}
ld calc(int h,int d,int a)
{
    if(a<need[h][d])return 0;
    if(a>=h+d)return 1;
    for(int i=1;i<=a;i++)
    {
        for(int j=1;j<=d;j++)
        {
            if(i>j)g[i-1][j-1]=get(h-1,d,a-i);
            else if(i<j)g[i-1][j-1]=get(h+1,d-j,a);
            else g[i-1][j-1]=get(h,d-j,a-i);
        }
    }
    ld z=simplex(a,d);
    return max(0.0L,min(1/z-1,1.0L));
}
void init(int md,int ma)
{
    int inf=ma+1;
    for(int i=0;i<=ma+1;i++)
    {
        for(int j=0;j<=md;j++)need[i][j]=inf;
    }
    for(int i=0;i<=md;i++)need[0][i]=0;
    for(int i=1;i<=ma+1;i++)need[i][0]=min(i,inf);
    for(int i=1;i<=md;i++)
    {
        for(int j=1;j<=ma;j++)
        {
            int mx=0;
            for(int k=1;k<=i;k++)
            {
                int mn=min(k+1+need[j-1][i],k+need[j][i-k]);
                if(k>=2)mn=min(mn,need[j+1][i-k]);
                mx=max(mx,min(mn,inf));
            }
            need[j][i]=mx;
        }
    }
    for(int i=2;i<=ma+md;i++)
    {
        for(int j=1;j<=ma;j++)
        {
            int d=i-j;
            if(d<1||d>md)continue;
            for(int k=1;k<=j;k++)f[k][d][j]=calc(k,d,j);
        }
    }
}
int main()
{
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
    int T;
    cin>>T;
    int md=0,ma=0;
    for(int i=0;i<T;i++)
    {
        cin>>q[i][0]>>q[i][1]>>q[i][2];
        md=max(md,q[i][1]);
        ma=max(ma,q[i][2]);
    }
    init(md,ma);
    cout<<fixed<<setprecision(6);
    for(int i=0;i<T;i++)cout<<get(q[i][0],q[i][1],q[i][2])<<'\n';
    return 0;
}