Files
hjha-server/PdkFriendServer/Logic/TrainingCollector.cs
xiaoou e5e1f4818e fix: V-learning 出牌率 0%→81% — MakePlay Type=None导致所有候选被当pass
根因: MakePlay(Type=CardType1.None)
  → play.Type==None 对所有候选为true
  → afterHand保持原样(未出牌)
  → 所有候选V值相同 → 始终选第一个(pass)

修复:
- play.Type→play.Ids.Length判断是否真pass
- EncodeAfterPlay更新ShenYuCard(出牌后剩余-1)
- EncodeAfterPlay设MaxPlayCard为出的牌
- isPass用Ids==null而非Type

效果: Play 117 Pass 28(81%出牌) vs ISMCTS -108/10局
2026-07-13 03:03:36 +08:00

167 lines
6.4 KiB
C#
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

using System;
using System.Collections.Generic;
using System.Linq;
using GameMessage;
using GameMessage.PaoDeKuaiF;
namespace PdkFriendServer.Logic
{
/// <summary>
/// 训练数据收集器 — 记录 (state, action, result) 三元组供 Python 训练。
/// 状态编码与 Python trainer 完全一致13rank×4suit×7ch=364 float。
/// </summary>
public static class TrainingCollector
{
public static bool Enabled = false;
private static readonly List<TrainingRecord> _records = new List<TrainingRecord>();
private static readonly int[] Ranks = { 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 16 };
private const int NumRanks = 13;
private const int NumSuits = 4;
private const int NumChannels = 7;
public const int InputDim = NumChannels * NumRanks * NumSuits; // 364
private struct TrainingRecord
{
public float[] State;
public int Action; // 0-12 = rank index, 13 = pass
public int Player;
}
public static void Record(PdkBotView view, CardTracker tracker, int actionIdx, int player)
{
if (!Enabled) return;
var state = EncodeState(view, tracker);
_records.Add(new TrainingRecord { State = state, Action = actionIdx, Player = player });
}
/// <summary>
/// 游戏结束,输出本轮训练数据到文件。
/// rewards: 每个玩家的奖励 (1=赢, 0=第二, -1=第三)
/// </summary>
public static void FlushGame(int gameId, float[] rewards)
{
if (!Enabled || _records.Count == 0) return;
var path = $"training_data/game_{gameId:D5}.csv";
System.IO.Directory.CreateDirectory("training_data");
using var w = new System.IO.StreamWriter(path);
w.WriteLine("state,action,reward");
foreach (var r in _records)
{
float reward = r.Player >= 0 && r.Player < rewards.Length ? rewards[r.Player] : 0;
var stateStr = string.Join(",", r.State.Select(f => f.ToString("F6")));
w.WriteLine($"\"{stateStr}\",{r.Action},{reward:F3}");
}
_records.Clear();
}
/// <summary>
/// 状态编码:与 Python encode_game_state 完全一致。
/// </summary>
/// <summary>
/// 编码"出牌后"状态:把我刚出的牌也标记为已见
/// </summary>
public static float[] EncodeAfterPlay(PdkBotView view, CardTracker tracker,
TCardInfoPdkF[] afterHand, int[] playedIds)
{
// 更新剩余张数:我打出了牌
int playedCount = playedIds?.Length ?? 0;
byte[] afterShenYu = new byte[view.ShenYuCard.Length];
Array.Copy(view.ShenYuCard, afterShenYu, view.ShenYuCard.Length);
afterShenYu[view.MyPos - 1] = (byte)Math.Max(0, afterShenYu[view.MyPos - 1] - playedCount);
var afterView = new PdkBotView
{
MyHand = afterHand,
MyPos = view.MyPos,
ShenYuCard = afterShenYu,
// 出牌后如果我不是passMaxPlayCard改成我出的牌
MaxPlayCard = playedIds != null && playedIds.Length > 0
? new PlayOutCardPdkF { Pos = (byte)view.MyPos,
GameNum = view.MyHand.FirstOrDefault(c => playedIds.Contains(c.ID)).GameNum,
Type = CardType1.DanZhang, Ids = playedIds.Select(id => (byte)id).ToArray() }
: view.MaxPlayCard,
Rule = view.Rule,
};
return EncodeState(afterView, tracker, playedIds);
}
public static float[] EncodeState(PdkBotView view, CardTracker tracker, int[] extraSeen = null)
{
var state = new float[InputDim];
// ch0: my hand
foreach (var c in view.MyHand)
{
if (c.GameState != 1) continue;
int ri = RankToIdx(c.GameNum);
int si = SuitToIdx(c.ID);
if (ri >= 0 && si >= 0)
state[ri * NumSuits + si] = 1.0f;
}
// ch1: seen cards (from tracker) + extraSeen
var pool = tracker.GetUnknownPool(view.MyHand);
var unseenIds = new HashSet<int>(pool.Select(c => (int)c.ID));
if (extraSeen != null)
foreach (var id in extraSeen) unseenIds.Remove(id);
for (int suit = 1; suit <= 4; suit++)
{
for (int n = 1; n <= 13; n++)
{
int id = (suit - 1) * 13 + n;
if (!unseenIds.Contains(id))
{
int g = n == 1 ? 14 : n == 2 ? 16 : n;
int ri = RankToIdx(g);
if (ri >= 0)
state[ChOffset(1) + ri * NumSuits + (suit - 1)] = 1.0f;
}
}
}
// ch2-ch6: scalar info
int myCount = view.MyHand.Count(c => c.GameState == 1);
state[ChOffset(2)] = (float)myCount / 16.0f;
int opp1 = view.MyPos % 3;
state[ChOffset(3)] = (float)view.ShenYuCard[opp1] / 16.0f;
int opp2 = (view.MyPos + 1) % 3;
state[ChOffset(4)] = (float)view.ShenYuCard[opp2] / 16.0f;
state[ChOffset(5) + view.MyPos % NumRanks] = 1.0f;
if (view.MaxPlayCard.GameNum > 0 && view.MaxPlayCard.Pos != view.MyPos)
{
int lr = RankToIdx(view.MaxPlayCard.GameNum);
if (lr >= 0)
state[ChOffset(6) + lr] = 1.0f;
}
return state;
}
public static int RankToIdx(int gameNum) => gameNum switch
{
3 => 0, 4 => 1, 5 => 2, 6 => 3, 7 => 4, 8 => 5, 9 => 6,
10 => 7, 11 => 8, 12 => 9, 13 => 10, 14 => 11, 16 => 12,
_ => -1
};
public static int ActionToIdx(PlayOutCardPdkF play)
{
if (play.Type == CardType1.None || play.Ids == null || play.Ids.Length == 0)
return 13; // pass
return RankToIdx(play.GameNum);
}
private static int SuitToIdx(int id)
{
int suit = (id - 1) / 13;
return suit >= 0 && suit < 4 ? suit : -1;
}
private static int ChOffset(int ch) => ch * NumRanks * NumSuits;
}
}