Skip to content

Commit 78413ae

Browse files
authored
Use NDArrayDescriptor.resolvingDynamicDimensions instead of modifying NDArrayDescriptor.shape directly (apple#74)
1 parent 7f9db7c commit 78413ae

1 file changed

Lines changed: 34 additions & 35 deletions

File tree

swift/Sources/CoreAILanguageModels/InferenceEngines/KVCache+CoreAI.swift

Lines changed: 34 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -241,14 +241,14 @@ struct StaticKVCache: CoreAIKVCache {
241241
self.currentCapacity = min(capacity ?? maxCapacityFromModel, maxCapacityFromModel)
242242
}
243243

244-
// Create modified requirements with adjusted sequence dimension.
245-
var keyReqsMod = keyReqs
246-
var valueReqsMod = valueReqs
247-
keyReqsMod.shape[seqDim] = self.currentCapacity
248-
valueReqsMod.shape[seqDim] = self.currentCapacity
244+
// Build concrete shapes with adjusted sequence dimension.
245+
var keyShape = keyReqs.shape
246+
var valueShape = valueReqs.shape
247+
keyShape[seqDim] = self.currentCapacity
248+
valueShape[seqDim] = self.currentCapacity
249249

250-
let keyResolved = keyReqsMod.resolvingDynamicDimensions(keyReqsMod.shape)
251-
let valueResolved = valueReqsMod.resolvingDynamicDimensions(valueReqsMod.shape)
250+
let keyResolved = keyReqs.resolvingDynamicDimensions(keyShape)
251+
let valueResolved = valueReqs.resolvingDynamicDimensions(valueShape)
252252

253253
let keyByteCount = keyResolved.minimumByteCount
254254
let valueByteCount = valueResolved.minimumByteCount
@@ -260,16 +260,16 @@ struct StaticKVCache: CoreAIKVCache {
260260
}
261261

262262
self.keyBinding = TensorBinding(
263-
metalBuffer: keyBuf, shape: keyReqsMod.shape,
263+
metalBuffer: keyBuf, shape: keyShape,
264264
strides: keyResolved.preferredStrides, scalarType: keyReqs.scalarType)
265265
self.valueBinding = TensorBinding(
266-
metalBuffer: valueBuf, shape: valueReqsMod.shape,
266+
metalBuffer: valueBuf, shape: valueShape,
267267
strides: valueResolved.preferredStrides, scalarType: valueReqs.scalarType)
268268

269269
// Log final allocation summary
270270
let fmt = ByteCountFormatter()
271271
fmt.countStyle = .memory
272-
let shapeDesc = KVCacheFactory.describeKVCacheStructure(shape: keyReqsMod.shape)
272+
let shapeDesc = KVCacheFactory.describeKVCacheStructure(shape: keyShape)
273273
CLILogger.log(
274274
"StaticKVCache allocated: \(shapeDesc), Total: \(fmt.string(fromByteCount: Int64(keyByteCount + valueByteCount)))"
275275
)
@@ -354,14 +354,14 @@ struct GrowingKVCache: CoreAIKVCache {
354354
self.maxCapacity = maxCapacityFromModel > 0 ? maxCapacityFromModel : Int.max
355355
self.currentCapacity = initialCapacity
356356

357-
// Create modified requirements with initial capacity.
358-
var keyReqsMod = keyReqs
359-
var valueReqsMod = valueReqs
360-
keyReqsMod.shape[sequenceDim] = self.currentCapacity
361-
valueReqsMod.shape[sequenceDim] = self.currentCapacity
357+
// Build concrete shapes with initial capacity.
358+
var keyShape = keyReqs.shape
359+
var valueShape = valueReqs.shape
360+
keyShape[sequenceDim] = self.currentCapacity
361+
valueShape[sequenceDim] = self.currentCapacity
362362

363-
let keyResolved = keyReqsMod.resolvingDynamicDimensions(keyReqsMod.shape)
364-
let valueResolved = valueReqsMod.resolvingDynamicDimensions(valueReqsMod.shape)
363+
let keyResolved = keyReqs.resolvingDynamicDimensions(keyShape)
364+
let valueResolved = valueReqs.resolvingDynamicDimensions(valueShape)
365365

366366
let keyByteCount = keyResolved.minimumByteCount
367367
let valueByteCount = valueResolved.minimumByteCount
@@ -373,16 +373,16 @@ struct GrowingKVCache: CoreAIKVCache {
373373
}
374374

375375
self.keyBinding = TensorBinding(
376-
metalBuffer: keyBuf, shape: keyReqsMod.shape,
376+
metalBuffer: keyBuf, shape: keyShape,
377377
strides: keyResolved.preferredStrides, scalarType: keyReqs.scalarType)
378378
self.valueBinding = TensorBinding(
379-
metalBuffer: valueBuf, shape: valueReqsMod.shape,
379+
metalBuffer: valueBuf, shape: valueShape,
380380
strides: valueResolved.preferredStrides, scalarType: valueReqs.scalarType)
381381

382382
// Log final allocation summary
383383
let fmt = ByteCountFormatter()
384384
fmt.countStyle = .memory
385-
let shapeDesc = KVCacheFactory.describeKVCacheStructure(shape: keyReqsMod.shape)
385+
let shapeDesc = KVCacheFactory.describeKVCacheStructure(shape: keyShape)
386386
CLILogger.log(
387387
"GrowingKVCache allocated (initial): \(shapeDesc), Total: \(fmt.string(fromByteCount: Int64(keyByteCount + valueByteCount)))"
388388
)
@@ -432,14 +432,14 @@ struct GrowingKVCache: CoreAIKVCache {
432432
}
433433
guard newCapacity > currentCapacity else { return nil }
434434

435-
// Create modified requirements with new capacity
436-
var keyReqsMod = keyReqsTemplate
437-
var valueReqsMod = valueReqsTemplate
438-
keyReqsMod.shape[sequenceDim] = newCapacity
439-
valueReqsMod.shape[sequenceDim] = newCapacity
435+
// Build concrete shapes with new capacity.
436+
var keyShape = keyReqsTemplate.shape
437+
var valueShape = valueReqsTemplate.shape
438+
keyShape[sequenceDim] = newCapacity
439+
valueShape[sequenceDim] = newCapacity
440440

441-
let keyResolved = keyReqsMod.resolvingDynamicDimensions(keyReqsMod.shape)
442-
let valueResolved = valueReqsMod.resolvingDynamicDimensions(valueReqsMod.shape)
441+
let keyResolved = keyReqsTemplate.resolvingDynamicDimensions(keyShape)
442+
let valueResolved = valueReqsTemplate.resolvingDynamicDimensions(valueShape)
443443

444444
let newKeyByteCount = keyResolved.minimumByteCount
445445
let newValueByteCount = valueResolved.minimumByteCount
@@ -455,11 +455,10 @@ struct GrowingKVCache: CoreAIKVCache {
455455
let oldValueBuf = valueBinding.metalBuffer
456456

457457
// Extract shape dimensions: [L, B, H, S, D]
458-
let shape = keyReqsMod.shape
459-
let l = shape[0]
460-
let b = shape[1]
461-
let h = shape[2]
462-
let d = shape[4]
458+
let l = keyShape[0]
459+
let b = keyShape[1]
460+
let h = keyShape[2]
461+
let d = keyShape[4]
463462
let oldS = currentCapacity
464463
let newS = newCapacity
465464
let mpsDataType = keyReqsTemplate.scalarType.mpsDataType
@@ -476,15 +475,15 @@ struct GrowingKVCache: CoreAIKVCache {
476475

477476
// Update bindings to new buffers (CPU metadata only — safe before GPU executes)
478477
keyBinding = TensorBinding(
479-
metalBuffer: newKeyBuf, shape: keyReqsMod.shape,
478+
metalBuffer: newKeyBuf, shape: keyShape,
480479
strides: keyResolved.preferredStrides, scalarType: keyReqsTemplate.scalarType)
481480
valueBinding = TensorBinding(
482-
metalBuffer: newValueBuf, shape: valueReqsMod.shape,
481+
metalBuffer: newValueBuf, shape: valueShape,
483482
strides: valueResolved.preferredStrides, scalarType: valueReqsTemplate.scalarType)
484483

485484
let fmt = ByteCountFormatter()
486485
fmt.countStyle = .memory
487-
let shapeDesc = KVCacheFactory.describeKVCacheStructure(shape: keyReqsMod.shape)
486+
let shapeDesc = KVCacheFactory.describeKVCacheStructure(shape: keyShape)
488487
CLILogger.log(
489488
"GrowingKVCache pipelined grow: \(currentCapacity)\(newCapacity), \(shapeDesc), Total: \(fmt.string(fromByteCount: Int64(newKeyByteCount + newValueByteCount)))"
490489
)

0 commit comments

Comments
 (0)