"""Optional homogeneous FP32 batching, with no increase in total padding.""" from dataclasses import replace from ._request import request_collator, predict_1k def _padded_tokens(entries, size): return sum(len(chunk) * max(row['input_tokens'] for row in chunk) for start in range(0, len(entries), size) for chunk in [entries[start:start + size]]) def _capacity(entries): if (len(entries) > 8 and len({row['kind'] for row in entries}) == 1 and _padded_tokens(entries, 32) <= _padded_tokens(entries, 8)): return 32 return 8 def predict_auto_1k(native, records): """Opt in to at most32 consecutive rows; same complete1024 FP32 contract. Default predict_1k remains unchanged. Mixed requests or extra padding retain B8. Every input admits first; errors never trigger partial outputs/retries. This changes physical shapes and can change FP32 rounding, not task semantics. """ records = list(records) if len(records) <= 8: return predict_1k(native, records, batch_size=8) guard = request_collator(native.collator, records) from decision_runtime import predict result = predict(replace(native, collator=guard), records, batch_size=_capacity(guard._encoded)) guard.finish() return result