如题。
#include <bits/stdc++.h>
#include <ext/pb_ds/priority_queue.hpp>
#define int long long
#define pq __gnu_pbds::priority_queue
#define pqg pq<int, vector<int>, greater<int> ()>
using namespace std;
struct ios {
inline char read() {
static const int inlen = 1 << 18 | 1;
static char buf[inlen], *s, *t;
return (s == t) && (t = (s = buf) + fread(buf, 1, inlen, stdin)), s == t ? -1 : *s++;
}
template<typename T> inline ios& operator>> (T &x) {
static char c11, boo;
for (c11 = read(), boo = 0; !isdigit(c11); c11 = read()) {
if (c11 == -1) return *this;
boo |= c11 == '-';
}
for (x = 0; isdigit(c11); c11 = read()) x = x * 10 + (c11 ^ '0');
boo && (x = -x);
return *this;
}
} fin;
struct exios {
template<typename _CharT, typename _Traits = char_traits<_CharT>>
struct typ {
typedef basic_ostream<_CharT, _Traits>& (* end) (basic_ostream<_CharT, _Traits>&);
};
template<typename T> friend exios &operator<<(exios &out, T num) {
if (num < 0) putchar('-'), num = -num;
if (num >= 10) out << num / 10;
putchar(num % 10 + '0');
return out;
}
friend exios &operator<<(exios &out, const char * s) { printf("%s", s); return out; }
friend exios &operator<<(exios &out, string s) { cout << s; return out; }
friend exios &operator<<(exios &out, typ<char>::end e) { puts(""); return out; }
} fout;
const int maxn = 1e6 + 10;
int n, k, rt;
vector<int> g[maxn];
int deg[maxn];
int dep[maxn];
pq<int, greater<int>> p[maxn];
void dfs(int u, int fa) {
dep[u] = dep[fa] + 1;
for (int v : g[u]) {
if (v == fa) continue;
dfs(v, u);
}
}
void solve(int u, int fa) {
if (deg[u] == 1) return;
for (int v : g[u]) {
if (v == fa) continue;
solve(v, u);
}
pq<int, greater<int>> res = pq<int, greater<int>>(), nw = pq<int, greater<int>>();
for (int v : g[u]) res.join(p[v]);
int x, y = res.top();
while (res.size() >= 2) {
x = res.top(), res.pop();
y = res.top();
if (x + y - dep[u] * 2 > k) nw.push(x);
}
nw.push(y);
p[u] = nw;
}
signed main() {
fin >> n >> k;
for (int i = 1, u, v; i < n; i++) {
fin >> u >> v;
g[u].push_back(v), g[v].push_back(u);
deg[u]++, deg[v]++;
}
for (int i = 1; i <= n; i++) if (deg[i] != 1) rt = i, i = n + 1;
dep[0] = -1;
dfs(rt, 0);
for (int i = 1; i <= n; i++) if (deg[i] == 1) p[i].push(dep[i]);
solve(rt, 0);
fout << p[rt].size() << endl;
return 0;
}