Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -184,7 +184,8 @@ A second snapshot, `InternDecisionModelStore.ensure(.showdown)`, is the same mod
battle actions ([FluidInference/intern-decision-0.8b-showdown-coreml](https://huggingface.co/FluidInference/intern-decision-0.8b-showdown-coreml)),
distilled from Intern-Decision-4B on a MacBook. It answers the same typed questions; given a battle state and the legal
moves and switches as `choice` options it matches its 4B teacher (24-6 vs poke-env's max-power player, 9-21 vs its
heuristic player over 30 battles). Buckets are 512, 640 and 1,024 tokens.
heuristic player over 30 battles). Buckets are 512, 640 and 1,024 tokens with int8 weights: about 0.85 GB in memory
with one bucket in use, 1.9 GB to download, the same 90 ms per decision as fp16.

## Demo

Expand Down
58 changes: 43 additions & 15 deletions Sources/FluidUse/InternDecision/InternDecisionModelStore.swift
Original file line number Diff line number Diff line change
Expand Up @@ -36,15 +36,17 @@ public enum InternDecisionModelStore {
let sha256: String
}

static let showdownRevision = "636dee2eacf95a18077e5c798245cc658d0f8747"
static let showdownRevision = "878e1404df1f257e02feef3a91e874fdccdd3780"

static func package(_ bucket: String, model: String, weights: String, manifest: String, config: String) -> [Asset] {
static func package(
_ bucket: String, name: String = "DecisionRow_fp16", model: String, weights: String, manifest: String,
config: String
) -> [Asset] {
[
Asset(path: "\(bucket)/config.json", sha256: config),
Asset(path: "\(bucket)/DecisionRow_fp16.mlpackage/Data/com.apple.CoreML/model.mlmodel", sha256: model),
Asset(
path: "\(bucket)/DecisionRow_fp16.mlpackage/Data/com.apple.CoreML/weights/weight.bin", sha256: weights),
Asset(path: "\(bucket)/DecisionRow_fp16.mlpackage/Manifest.json", sha256: manifest),
Asset(path: "\(bucket)/\(name).mlpackage/Data/com.apple.CoreML/model.mlmodel", sha256: model),
Asset(path: "\(bucket)/\(name).mlpackage/Data/com.apple.CoreML/weights/weight.bin", sha256: weights),
Asset(path: "\(bucket)/\(name).mlpackage/Manifest.json", sha256: manifest),
]
}

Expand All @@ -70,26 +72,31 @@ public enum InternDecisionModelStore {
manifest: "363c3adc8aea3a5a70679e78f062ffa5b8651e14c9ec759d556fc7976f3bfd35",
config: "2fa46a5ee9844b89eee0f43c04206d544540a4fe7113cd74229266b546030b69")

/// The int8 (weight-only, per-channel) buckets: same speed as fp16, about 0.85 GB in memory with one bucket in
/// use instead of about 1.5 GB, and no measurable change in play.
static let showdownAssets: [Asset] =
[
Asset(path: "config.json", sha256: "9f41af5eb81d91e736045849e0b5cc65948edd53a7893043ad391fefeddc89f7"),
Asset(path: "embeddings.f16", sha256: "703d76a6923d2b1fad57d11e9e11dc60b2c2db6b448ab4a1a98aa7e1c6327877"),
Asset(path: "tokenizer.json", sha256: "94a639c4b33b192cc5a22cd3d7f0aaf6d97efa9577957a0382b55292ca4f0f00"),
]
+ package(
"L512_F8", model: "e090f2632d7c6c944c1f102591a81a4836c7bf3a55a3b925c652cbac180e8968",
weights: "c0dff14d537bff0193e81f61da39e6edcb8a6ded4d6c7cf6f48c9a5128860875",
manifest: "8eaf456ff4f5ff1916f93354ac64521fd7e2fb0a0aed1a320cec799afdde1de2",
"L512_F8", name: "DecisionRow_w8",
model: "1c53fecc9fc67d03884a593448074113e1db040b9908afe59e6482d67381b290",
weights: "4433b38de2fda23446174fb0d011736c195c32dd9a924c2b8c53ad028b50c21a",
manifest: "171b2dca62328c5358becfa36d43c186bd64678ba4fed2b1d14ff8f21af02da5",
config: "544bc212531abb4ebd47fbbcb5b6cba3084186115971d84eaf874e736a4490bc")
+ package(
"L640_F8", model: "29f4535b956e1999aef623a02f73ba42f7a320f1ca33ae30446895b160f0afb9",
weights: "3aebedd7762bd1230f491dfe8d1fbd8d2c77785cdac7836f907f36e721bce65b",
manifest: "ea6d9f587d2d1d81b28e20cdf3ad1fa0b47ae1f14bc0053dac0dabeea3f6e8d3",
"L640_F8", name: "DecisionRow_w8",
model: "2ed49dcf8978cbed9dabc19365cba6507eca638eb03baab0f977a61e39d4b91f",
weights: "3711e818b4612c7024a543fbc6bf0b22acf43b647e2930b2ce624f70dcb8cfcc",
manifest: "759db0f58151fbcd5c27fbb80f918bb1ae419e925a04de42fd03e6c82419685f",
config: "b6301674536dc7ead157ad70b5d7b5990aad7c7f5073379638ae71dc775c6ef9")
+ package(
"L1024_F16", model: "1129b6bd92a4d4e368968c81b26d699ce51fb9f39e0844bb3790a239fab7f133",
weights: "2fbffc2bc4be9b6d965e2a1ca55bf64e49e603f74470b29f12c3b1589dedad80",
manifest: "3e730b08ddaffa138a77647724cecc6895f444f76c4ae36328c719f28d52aad4",
"L1024_F16", name: "DecisionRow_w8",
model: "4ce919d28e3ce2661b3e04164e44cfa039badfdc34c22eba1384fd2bb6df1fce",
weights: "db8c66f760c48627e7fe4957c82c9b5ff0ea4ee23a6da9f570b138f27dc8e006",
manifest: "add62157ac7623bc8cde8820c55eafd23322a4fb815f69a451685bb9afee9461",
config: "b0b3a036af825b3457e4a0f5225648021def8513d44d7a77b983ab55add6bc2f")

/// Ensure the snapshot exists in the FluidUse cache and return its directory. Files are checksummed once per
Expand Down Expand Up @@ -140,10 +147,31 @@ public enum InternDecisionModelStore {
}
progress?(asset.path, size)
}
try Self.removeStalePackages(in: directory, keeping: model.assets)
try Data().write(to: verified)
return directory
}

/// After a revision change, packages (and their compiled forms) that the pinned snapshot no longer lists must
/// go, or `InternDecisionManager.load` could pick an older precision left in the same folder.
static func removeStalePackages(in directory: URL, keeping assets: [Asset]) throws {
let manager = FileManager.default
let kept = Set(
assets.compactMap { asset -> String? in
guard let range = asset.path.range(of: ".mlpackage/") else { return nil }
return String(asset.path[..<range.lowerBound])
})
guard let files = manager.enumerator(at: directory, includingPropertiesForKeys: nil) else { return }
var stale: [URL] = []
for case let url as URL in files where ["mlpackage", "mlmodelc"].contains(url.pathExtension) {
let relative = url.path.replacingOccurrences(of: directory.path + "/", with: "")
let stem = String(relative.dropLast(url.pathExtension.count + 1))
if !kept.contains(stem) { stale.append(url) }
files.skipDescendants()
}
for url in stale { try manager.removeItem(at: url) }
}

private static func checksum(of file: URL) throws -> String {
let handle = try FileHandle(forReadingFrom: file)
defer { try? handle.close() }
Expand Down
Loading